149 lines
5.9 KiB
Python
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())
|