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