Files
leocrm/app/plugins/builtins/unified_search/query_understanding.py
T
Agent Zero 5c3fc027bc fix: strip markdown code fences from LLM responses before json.loads
- glm-5.2 returns JSON wrapped in ```json ... ``` code fences
- Added fence stripping in services.py, query_understanding.py, jobs.py
2026-07-19 02:30:07 +02:00

169 lines
5.7 KiB
Python

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