165 lines
7.2 KiB
Python
165 lines
7.2 KiB
Python
from __future__ import annotations
|
|
|
|
import argparse
|
|
import json
|
|
import urllib.request
|
|
from pathlib import Path
|
|
from typing import Any
|
|
|
|
from rag_embedding_providers import LOCAL_HASHING_MODEL, LOCAL_HASHING_PROVIDER, embed_texts, provider_metadata
|
|
|
|
|
|
DEFAULT_ADAPTER_URL = "http://docker-gpu.cin.su:8011/rpc"
|
|
|
|
|
|
def adapter_call(adapter_url: str, method: str, payload: dict[str, Any], *, timeout_seconds: int = 180) -> dict[str, Any]:
|
|
body = json.dumps({"method": method, "payload": payload}, ensure_ascii=False).encode("utf-8")
|
|
request = urllib.request.Request(adapter_url, data=body, headers={"Content-Type": "application/json"}, method="POST")
|
|
with urllib.request.urlopen(request, timeout=timeout_seconds) as response:
|
|
result = json.loads(response.read().decode("utf-8"))
|
|
if not isinstance(result, dict):
|
|
raise ValueError(f"{method} returned a non-object response")
|
|
return result
|
|
|
|
|
|
def batched(items: list[dict[str, Any]], size: int) -> list[list[dict[str, Any]]]:
|
|
return [items[index : index + size] for index in range(0, len(items), size)]
|
|
|
|
|
|
def embedding_model_label(*, provider: str, model: str) -> str:
|
|
normalized = provider_metadata(provider=provider, model=model, dimensions=1).get("embedding_provider")
|
|
return f"{normalized}:{model}" if normalized != "local_hashing" else model
|
|
|
|
|
|
def embed_pending_semantic_cache(
|
|
*,
|
|
adapter_url: str,
|
|
base_id: str,
|
|
kind: str = "",
|
|
limit: int = 100,
|
|
batch_size: int = 16,
|
|
embedding_provider: str = LOCAL_HASHING_PROVIDER,
|
|
embedding_model: str = LOCAL_HASHING_MODEL,
|
|
dimensions: int = 384,
|
|
embedding_base_url: str = "",
|
|
embedding_api_key_env: str = "OPENAI_API_KEY",
|
|
dry_run: bool = False,
|
|
timeout_seconds: int = 180,
|
|
) -> dict[str, Any]:
|
|
pending_payload: dict[str, Any] = {"base_id": base_id, "limit": limit, "vector_status": "pending_embedding"}
|
|
if kind:
|
|
pending_payload["kind"] = kind
|
|
pending = adapter_call(adapter_url, "semantic.cache.pending", pending_payload, timeout_seconds=timeout_seconds)
|
|
if pending.get("status") != "ok":
|
|
return {"status": pending.get("status") or "error", "error": pending.get("error"), "pending": pending}
|
|
documents = [item for item in pending.get("documents") or [] if isinstance(item, dict)]
|
|
model_label = embedding_model_label(provider=embedding_provider, model=embedding_model)
|
|
upserts: list[dict[str, Any]] = []
|
|
skipped: list[dict[str, Any]] = []
|
|
for batch in batched(documents, max(int(batch_size or 1), 1)):
|
|
texts = [str(item.get("text") or "") for item in batch]
|
|
vectors = embed_texts(
|
|
texts,
|
|
provider=embedding_provider,
|
|
model=embedding_model,
|
|
dimensions=dimensions,
|
|
base_url=embedding_base_url,
|
|
api_key_env=embedding_api_key_env,
|
|
)
|
|
for item, vector in zip(batch, vectors):
|
|
document_id = str(item.get("document_id") or "")
|
|
content_sha1 = str(item.get("content_sha1") or "")
|
|
if not document_id or not content_sha1 or not vector:
|
|
skipped.append({"document_id": document_id or None, "reason": "missing_document_id_content_sha1_or_embedding"})
|
|
continue
|
|
upsert_payload = {
|
|
"base_id": base_id,
|
|
"document_id": document_id,
|
|
"content_sha1": content_sha1,
|
|
"embedding_model": model_label,
|
|
"embedding": vector,
|
|
}
|
|
if dry_run:
|
|
upserts.append({"status": "dry_run", "document_id": document_id, "content_sha1": content_sha1, "dimensions": len(vector)})
|
|
continue
|
|
result = adapter_call(adapter_url, "semantic.cache.embedding.upsert", upsert_payload, timeout_seconds=timeout_seconds)
|
|
upserts.append(
|
|
{
|
|
"status": result.get("status"),
|
|
"error": result.get("error"),
|
|
"document_id": document_id,
|
|
"content_sha1": content_sha1,
|
|
"dimensions": result.get("dimensions") or len(vector),
|
|
}
|
|
)
|
|
return {
|
|
"schema": "onec_semantic_cache_embedding_worker.v1",
|
|
"status": "ok",
|
|
"base_id": base_id,
|
|
"adapter_url": adapter_url,
|
|
"dry_run": bool(dry_run),
|
|
"embedding": {
|
|
"provider": provider_metadata(provider=embedding_provider, model=embedding_model, dimensions=dimensions, base_url=embedding_base_url).get("embedding_provider"),
|
|
"model": embedding_model,
|
|
"stored_embedding_model": model_label,
|
|
"dimensions": dimensions,
|
|
},
|
|
"counts": {
|
|
"pending": len(documents),
|
|
"processed": len(upserts),
|
|
"stored": len([item for item in upserts if item.get("status") == "ok"]),
|
|
"conflicts": len([item for item in upserts if item.get("status") == "conflict"]),
|
|
"skipped": len(skipped),
|
|
"errors": len([item for item in upserts if item.get("status") not in {"ok", "dry_run", "conflict"}]),
|
|
},
|
|
"upserts": upserts,
|
|
**({"skipped": skipped} if skipped else {}),
|
|
}
|
|
|
|
|
|
def main() -> int:
|
|
parser = argparse.ArgumentParser(description="Embed pending documents from the 1C adapter semantic cache.")
|
|
parser.add_argument("--adapter-url", default=DEFAULT_ADAPTER_URL)
|
|
parser.add_argument("--base-id", required=True)
|
|
parser.add_argument("--kind", default="")
|
|
parser.add_argument("--limit", type=int, default=100)
|
|
parser.add_argument("--batch-size", type=int, default=16)
|
|
parser.add_argument("--embedding-provider", default=LOCAL_HASHING_PROVIDER, choices=[LOCAL_HASHING_PROVIDER, "openai-compatible"])
|
|
parser.add_argument("--embedding-model", default=LOCAL_HASHING_MODEL)
|
|
parser.add_argument("--dimensions", type=int, default=384)
|
|
parser.add_argument("--embedding-base-url", default="")
|
|
parser.add_argument("--embedding-api-key-env", default="OPENAI_API_KEY")
|
|
parser.add_argument("--timeout-seconds", type=int, default=180)
|
|
parser.add_argument("--dry-run", action="store_true")
|
|
parser.add_argument("--json", action="store_true")
|
|
args = parser.parse_args()
|
|
|
|
result = embed_pending_semantic_cache(
|
|
adapter_url=args.adapter_url,
|
|
base_id=args.base_id,
|
|
kind=args.kind,
|
|
limit=args.limit,
|
|
batch_size=args.batch_size,
|
|
embedding_provider=args.embedding_provider,
|
|
embedding_model=args.embedding_model,
|
|
dimensions=args.dimensions,
|
|
embedding_base_url=args.embedding_base_url,
|
|
embedding_api_key_env=args.embedding_api_key_env,
|
|
dry_run=args.dry_run,
|
|
timeout_seconds=args.timeout_seconds,
|
|
)
|
|
if args.json:
|
|
print(json.dumps(result, ensure_ascii=False, indent=2))
|
|
else:
|
|
counts = result.get("counts") or {}
|
|
print(
|
|
"semantic cache embeddings: "
|
|
f"pending={counts.get('pending')} processed={counts.get('processed')} "
|
|
f"stored={counts.get('stored')} conflicts={counts.get('conflicts')} errors={counts.get('errors')}"
|
|
)
|
|
return 0 if result.get("status") == "ok" and int((result.get("counts") or {}).get("errors") or 0) == 0 else 1
|
|
|
|
|
|
if __name__ == "__main__":
|
|
raise SystemExit(main())
|