Files
llm/scripts/search_1c_semantic_cache.py
T

149 lines
5.9 KiB
Python

from __future__ import annotations
import argparse
import json
from typing import Any
from embed_1c_semantic_cache import DEFAULT_ADAPTER_URL, adapter_call, embed_pending_semantic_cache
from rag_embedding_providers import LOCAL_HASHING_MODEL, LOCAL_HASHING_PROVIDER, embed_texts, provider_metadata
def search_semantic_cache(
*,
adapter_url: str,
base_id: str,
query: str,
kind: str = "",
limit: int = 20,
scan_limit: int = 1000,
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",
validate_candidates: bool = False,
validation_limit: int | None = None,
embed_pending: bool = False,
embed_limit: int = 100,
embed_batch_size: int = 16,
timeout_seconds: int = 180,
) -> dict[str, Any]:
embed_result = None
if embed_pending:
embed_result = embed_pending_semantic_cache(
adapter_url=adapter_url,
base_id=base_id,
kind=kind,
limit=embed_limit,
batch_size=embed_batch_size,
embedding_provider=embedding_provider,
embedding_model=embedding_model,
dimensions=dimensions,
embedding_base_url=embedding_base_url,
embedding_api_key_env=embedding_api_key_env,
timeout_seconds=timeout_seconds,
)
query_embedding = embed_texts(
[query],
provider=embedding_provider,
model=embedding_model,
dimensions=dimensions,
base_url=embedding_base_url,
api_key_env=embedding_api_key_env,
)[0]
payload: dict[str, Any] = {
"base_id": base_id,
"query": query,
"query_embedding": query_embedding,
"limit": limit,
"scan_limit": scan_limit,
"validate_candidates": validate_candidates,
}
if kind:
payload["kind"] = kind
if validation_limit is not None:
payload["validation_limit"] = validation_limit
result = adapter_call(adapter_url, "semantic.cache.search", payload, timeout_seconds=timeout_seconds)
result.setdefault("client_embedding", {})
result["client_embedding"] = {
"provider": provider_metadata(provider=embedding_provider, model=embedding_model, dimensions=len(query_embedding), base_url=embedding_base_url).get("embedding_provider"),
"model": embedding_model,
"dimensions": len(query_embedding),
}
if embed_result is not None:
result["embedding_refresh"] = {
"status": embed_result.get("status"),
"counts": embed_result.get("counts") or {},
"embedding": embed_result.get("embedding") or {},
}
return result
def print_match(match: dict[str, Any], position: int) -> None:
obj = match.get("object") or {}
cache = match.get("cache") or {}
print(f"{position}. score={float(match.get('score') or 0):.4f} match_by={match.get('match_by')}")
print(f" object={obj.get('kind')} {obj.get('name') or obj.get('guid')}")
print(f" document_id={match.get('document_id')}")
print(f" embedding_model={cache.get('embedding_model')} vector_status={cache.get('vector_status')}")
text = str(match.get("text_preview") or "").replace("\n", " ").strip()
if len(text) > 500:
text = text[:497].rstrip() + "..."
print(f" text={text}")
def main() -> int:
parser = argparse.ArgumentParser(description="Search the 1C adapter semantic cache with a computed query embedding.")
parser.add_argument("query")
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=20)
parser.add_argument("--scan-limit", type=int, default=1000)
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("--validate-candidates", action="store_true")
parser.add_argument("--validation-limit", type=int)
parser.add_argument("--embed-pending", action="store_true")
parser.add_argument("--embed-limit", type=int, default=100)
parser.add_argument("--embed-batch-size", type=int, default=16)
parser.add_argument("--timeout-seconds", type=int, default=180)
parser.add_argument("--json", action="store_true")
args = parser.parse_args()
result = search_semantic_cache(
adapter_url=args.adapter_url,
base_id=args.base_id,
query=args.query,
kind=args.kind,
limit=args.limit,
scan_limit=args.scan_limit,
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,
validate_candidates=args.validate_candidates,
validation_limit=args.validation_limit,
embed_pending=args.embed_pending,
embed_limit=args.embed_limit,
embed_batch_size=args.embed_batch_size,
timeout_seconds=args.timeout_seconds,
)
if args.json:
print(json.dumps(result, ensure_ascii=False, indent=2))
elif result.get("matches"):
print(f"semantic_cache_status={result.get('status')} matches={len(result.get('matches') or [])}")
for position, match in enumerate(result.get("matches") or [], start=1):
print_match(match, position)
else:
print(f"No matches. status={result.get('status')} error={result.get('error')}")
return 0 if result.get("status") == "ok" else 1
if __name__ == "__main__":
raise SystemExit(main())