Files
llm/scripts/embed_1c_semantic_cache.py
T

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())