"""KI query understanding and result aggregation via LiteLLM.""" from __future__ import annotations import os import json import logging import uuid from typing import Any import litellm from sqlalchemy.ext.asyncio import AsyncSession logger = logging.getLogger(__name__) DEFAULT_LLM_MODEL = os.environ.get('SEARCH_LLM_MODEL', 'ollama/deepseek-v4-flash') QUERY_ANALYZE_SYSTEM = ( "Du bist ein Query-Analyzer fuer ein CRM. " "Analysiere die Suchanfrage und gib JSON zurueck: " '{"normalized_query": str, "entities": {"person": str|null, "company": str|null, "topic": str|null}, ' '"intent": str, "semantic_terms": [str], "suggested_filters": {}}' ) RESULT_AGGREGATE_SYSTEM = ( "Du bist ein Result-Aggregator. Fasse Ergebnisse zusammen und generiere Facetten: " '{"summary": str, "facets": {"types": {}, "dates": {}, "people": []}, "suggestions": [str]}' ) async def _get_api_credentials( db: AsyncSession | None, tenant_id: uuid.UUID | None ) -> tuple[str | None, str | None, str | None]: """Get API key, base_url and provider_type from the default AI provider in DB. Falls back to API_KEY_OLLAMA_CLOUD env var. Returns (api_key, base_url, provider_type). """ if db and tenant_id: try: from app.plugins.builtins.ai_assistant.services import get_default_provider provider = await get_default_provider(db, tenant_id) if provider and provider.api_key: return provider.api_key, provider.base_url, provider.provider_type except Exception: logger.debug("Failed to get provider from DB, falling back to env") env_key = os.environ.get('API_KEY_OLLAMA_CLOUD', '') return (env_key if env_key else None), None, None def _build_model(model: str, provider_type: str | None) -> str: """Build litellm model string with provider prefix.""" if provider_type: model_parts = model.split("/", 1) return f"{provider_type}/{model_parts[-1]}" return model def _fallback_query_analysis(query: str) -> dict[str, Any]: return { "normalized_query": query, "entities": {}, "intent": "search", "semantic_terms": [], "suggested_filters": {}, } def _fallback_aggregate(results: list[dict], query: str) -> dict[str, Any]: return { "summary": f"{len(results)} Ergebnisse gefunden", "facets": {}, "suggestions": [], } async def llm_analyze_query( query: str, db: AsyncSession | None = None, tenant_id: uuid.UUID | None = None, ) -> dict[str, Any]: """Analyze a search query using LLM for intent, entities, and semantic terms. Falls back to a simple dict if LLM fails. """ try: api_key, api_base, provider_type = await _get_api_credentials(db, tenant_id) model = _build_model(DEFAULT_LLM_MODEL, provider_type) litellm_kwargs: dict[str, Any] = dict( model=model, messages=[ {"role": "system", "content": QUERY_ANALYZE_SYSTEM}, {"role": "user", "content": query}, ], temperature=0.1, max_tokens=500, response_format={"type": "json_object"}, ) if api_key: litellm_kwargs["api_key"] = api_key if api_base: litellm_kwargs["api_base"] = api_base response = await litellm.acompletion(**litellm_kwargs) content = response.choices[0].message.content # Strip markdown code fences if present content = content.strip() if content.startswith("```"): content = content.split("\n", 1)[-1] if "\n" in content else content[3:] if content.endswith("```"): content = content[:-3].strip() return json.loads(content) except Exception: logger.warning("LLM query analysis failed, using fallback", exc_info=True) return _fallback_query_analysis(query) async def llm_aggregate_results( results: list[dict], query: str, db: AsyncSession | None = None, tenant_id: uuid.UUID | None = None, ) -> dict[str, Any]: """Aggregate search results using LLM for summary, facets, and suggestions. Falls back to a simple dict if LLM fails. """ if not results: return _fallback_aggregate(results, query) try: api_key, api_base, provider_type = await _get_api_credentials(db, tenant_id) model = _build_model(DEFAULT_LLM_MODEL, provider_type) # Truncate results to avoid token overflow compact = [ {"entity_type": r.get("entity_type"), "title": r.get("title", "")[:100]} for r in results[:50] ] user_msg = json.dumps({"query": query, "results": compact}) litellm_kwargs: dict[str, Any] = dict( model=model, messages=[ {"role": "system", "content": RESULT_AGGREGATE_SYSTEM}, {"role": "user", "content": user_msg}, ], temperature=0.1, max_tokens=1000, response_format={"type": "json_object"}, ) if api_key: litellm_kwargs["api_key"] = api_key if api_base: litellm_kwargs["api_base"] = api_base response = await litellm.acompletion(**litellm_kwargs) content = response.choices[0].message.content # Strip markdown code fences if present content = content.strip() if content.startswith("```"): content = content.split("\n", 1)[-1] if "\n" in content else content[3:] if content.endswith("```"): content = content[:-3].strip() return json.loads(content) except Exception: logger.warning("LLM result aggregation failed, using fallback", exc_info=True) return _fallback_aggregate(results, query)