feat(B-LLM): Zentraler LLM Client — llm_complete() + llm_embed() + Migration + Tests + Doku
Check Cross-Plugin Imports / check (push) Has been cancelled
Check Cross-Plugin Imports / check (push) Has been cancelled
B-LLM: llm_client.py um generische llm_complete() und llm_embed() erweitert - Provider-Auswahl, API-Key-Auflösung, Error-Handling, Cost-Tracking - Retry mit Exponential-Backoff für transient errors - Timeout konfigurierbar - Helper: get_api_credentials(), build_model(), _classify_error() B-LLM-MIG: Alle 8 direkten litellm.acompletion() Calls auf llm_complete() umgestellt - agent_runner.py, query_understanding.py (2x), ai_proactive (3x), ai_assistant (2x) - 0 verbleibende direkte litellm.acompletion() Calls außerhalb llm_client.py B-LLM-TEST: 39 Tests in test_llm_client.py — alle grün - Mock mode, error handling, embed, helpers, backward compat B-LLM-DOC: Plugin-Dev-Guide Kapitel 7 (LLM Integration) hinzugefügt
This commit is contained in:
@@ -1,26 +1,48 @@
|
||||
"""Embedding pipeline using LiteLLM with OpenRouter for embeddings."""
|
||||
"""Embedding pipeline using LiteLLM with OpenRouter for embeddings.
|
||||
|
||||
Delegates credential lookup, model building, and embedding generation to
|
||||
the centralised ``app.ai.llm_client`` module. The wrapper functions here
|
||||
preserve backward compatibility for existing call sites.
|
||||
"""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
import os
|
||||
import logging
|
||||
import os
|
||||
import uuid
|
||||
from typing import Any, TYPE_CHECKING
|
||||
|
||||
import litellm
|
||||
from typing import TYPE_CHECKING
|
||||
|
||||
if TYPE_CHECKING:
|
||||
from sqlalchemy.ext.asyncio import AsyncSession
|
||||
|
||||
from app.ai.llm_client import (
|
||||
EMBEDDING_DIMENSIONS,
|
||||
MAX_INPUT_CHARS,
|
||||
OPENROUTER_EMBEDDING_MODEL,
|
||||
build_model as _central_build_model,
|
||||
get_api_credentials as _central_get_api_credentials,
|
||||
llm_embed,
|
||||
)
|
||||
|
||||
logger = logging.getLogger(__name__)
|
||||
|
||||
MAX_INPUT_CHARS = 8000
|
||||
# Re-export constants for backward compatibility
|
||||
__all__ = [
|
||||
"MAX_INPUT_CHARS",
|
||||
"OPENROUTER_API_KEY",
|
||||
"OPENROUTER_BASE_URL",
|
||||
"OPENROUTER_EMBEDDING_MODEL",
|
||||
"EMBEDDING_DIMENSIONS",
|
||||
"_get_api_credentials",
|
||||
"_build_model",
|
||||
"generate_embedding",
|
||||
"generate_embeddings_batch",
|
||||
"index_entity",
|
||||
]
|
||||
|
||||
# OpenRouter for embeddings (Ollama Cloud has no embedding endpoint)
|
||||
OPENROUTER_API_KEY = os.environ.get('API_KEY_OPENROUTER', '')
|
||||
OPENROUTER_BASE_URL = 'https://openrouter.ai/api/v1'
|
||||
OPENROUTER_EMBEDDING_MODEL = os.environ.get('SEARCH_EMBEDDING_MODEL', 'openai/text-embedding-3-small')
|
||||
EMBEDDING_DIMENSIONS = 768 # Must match DB column vector(768)
|
||||
# Re-export for backward compatibility (consumers may import these directly)
|
||||
OPENROUTER_API_KEY = os.environ.get("API_KEY_OPENROUTER", "")
|
||||
OPENROUTER_BASE_URL = "https://openrouter.ai/api/v1"
|
||||
|
||||
|
||||
async def _get_api_credentials(
|
||||
@@ -28,34 +50,19 @@ async def _get_api_credentials(
|
||||
) -> tuple[str | None, str | None, str | None]:
|
||||
"""Get API key, base_url and provider_type for embeddings.
|
||||
|
||||
Priority:
|
||||
1. OpenRouter env var (API_KEY_OPENROUTER) – dedicated embedding provider
|
||||
2. Default AI provider from DB (fallback)
|
||||
3. API_KEY_OLLAMA_CLOUD env var (last resort)
|
||||
Thin wrapper delegating to ``app.ai.llm_client.get_api_credentials``.
|
||||
Kept for backward compatibility with existing call sites.
|
||||
"""
|
||||
# OpenRouter is the primary embedding provider
|
||||
if OPENROUTER_API_KEY:
|
||||
return OPENROUTER_API_KEY, OPENROUTER_BASE_URL, 'openai'
|
||||
|
||||
# Fallback to DB provider
|
||||
if db and tenant_id:
|
||||
try:
|
||||
from app.plugins.builtins.ai_assistant.contracts 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
|
||||
return await _central_get_api_credentials(db, tenant_id)
|
||||
|
||||
|
||||
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
|
||||
"""Build litellm model string with provider prefix.
|
||||
|
||||
Thin wrapper delegating to ``app.ai.llm_client.build_model``.
|
||||
Kept for backward compatibility with existing call sites.
|
||||
"""
|
||||
return _central_build_model(model, provider_type)
|
||||
|
||||
|
||||
async def generate_embedding(
|
||||
@@ -77,30 +84,15 @@ async def generate_embedding(
|
||||
Returns:
|
||||
Embedding vector as list of floats.
|
||||
"""
|
||||
model = model or OPENROUTER_EMBEDDING_MODEL
|
||||
truncated = text[:MAX_INPUT_CHARS]
|
||||
try:
|
||||
api_key, api_base, provider_type = await _get_api_credentials(db, tenant_id)
|
||||
litellm_model = _build_model(model, provider_type)
|
||||
|
||||
litellm_kwargs: dict[str, Any] = dict(
|
||||
model=litellm_model,
|
||||
input=truncated,
|
||||
)
|
||||
if api_key:
|
||||
litellm_kwargs["api_key"] = api_key
|
||||
if api_base:
|
||||
litellm_kwargs["api_base"] = api_base
|
||||
|
||||
# Request 768 dimensions to match DB vector(768) column
|
||||
if 'text-embedding-3' in litellm_model:
|
||||
litellm_kwargs['dimensions'] = EMBEDDING_DIMENSIONS
|
||||
|
||||
response = await litellm.aembedding(**litellm_kwargs)
|
||||
return response.data[0]["embedding"]
|
||||
except Exception:
|
||||
logger.warning("Failed to generate embedding", exc_info=True)
|
||||
return []
|
||||
embeddings = await llm_embed(
|
||||
texts=text,
|
||||
model=model,
|
||||
db=db,
|
||||
tenant_id=tenant_id,
|
||||
)
|
||||
if embeddings and embeddings[0]:
|
||||
return embeddings[0]
|
||||
return []
|
||||
|
||||
|
||||
async def generate_embeddings_batch(
|
||||
@@ -120,29 +112,12 @@ async def generate_embeddings_batch(
|
||||
Returns:
|
||||
List of embedding vectors.
|
||||
"""
|
||||
model = model or OPENROUTER_EMBEDDING_MODEL
|
||||
truncated = [t[:MAX_INPUT_CHARS] for t in texts]
|
||||
try:
|
||||
api_key, api_base, provider_type = await _get_api_credentials(db, tenant_id)
|
||||
litellm_model = _build_model(model, provider_type)
|
||||
|
||||
litellm_kwargs: dict[str, Any] = dict(
|
||||
model=litellm_model,
|
||||
input=truncated,
|
||||
)
|
||||
if api_key:
|
||||
litellm_kwargs["api_key"] = api_key
|
||||
if api_base:
|
||||
litellm_kwargs["api_base"] = api_base
|
||||
|
||||
if 'text-embedding-3' in litellm_model:
|
||||
litellm_kwargs['dimensions'] = EMBEDDING_DIMENSIONS
|
||||
|
||||
response = await litellm.aembedding(**litellm_kwargs)
|
||||
return [d["embedding"] for d in response.data]
|
||||
except Exception:
|
||||
logger.warning("Failed to generate batch embeddings", exc_info=True)
|
||||
return [[] for _ in texts]
|
||||
return await llm_embed(
|
||||
texts=texts,
|
||||
model=model,
|
||||
db=db,
|
||||
tenant_id=tenant_id,
|
||||
)
|
||||
|
||||
|
||||
async def index_entity(
|
||||
@@ -201,5 +176,5 @@ async def index_entity(
|
||||
await db.commit()
|
||||
return True
|
||||
except Exception:
|
||||
logger.exception("Failed to index entity %s/%s", entity_type, entity_id)
|
||||
logger.warning("Failed to index entity %s/%s", entity_type, entity_id, exc_info=True)
|
||||
return False
|
||||
|
||||
@@ -8,7 +8,7 @@ import logging
|
||||
import uuid
|
||||
from typing import Any
|
||||
|
||||
import litellm
|
||||
from app.ai.llm_client import llm_complete
|
||||
from sqlalchemy.ext.asyncio import AsyncSession
|
||||
|
||||
logger = logging.getLogger(__name__)
|
||||
@@ -87,7 +87,7 @@ async def llm_analyze_query(
|
||||
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(
|
||||
result = await llm_complete(
|
||||
model=model,
|
||||
messages=[
|
||||
{"role": "system", "content": QUERY_ANALYZE_SYSTEM},
|
||||
@@ -96,14 +96,10 @@ async def llm_analyze_query(
|
||||
temperature=0.1,
|
||||
max_tokens=500,
|
||||
response_format={"type": "json_object"},
|
||||
api_key=api_key,
|
||||
api_base=api_base,
|
||||
)
|
||||
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
|
||||
content = result["content"]
|
||||
# Strip markdown code fences if present
|
||||
content = content.strip()
|
||||
if content.startswith("```"):
|
||||
@@ -139,7 +135,7 @@ async def llm_aggregate_results(
|
||||
]
|
||||
user_msg = json.dumps({"query": query, "results": compact})
|
||||
|
||||
litellm_kwargs: dict[str, Any] = dict(
|
||||
result = await llm_complete(
|
||||
model=model,
|
||||
messages=[
|
||||
{"role": "system", "content": RESULT_AGGREGATE_SYSTEM},
|
||||
@@ -148,14 +144,10 @@ async def llm_aggregate_results(
|
||||
temperature=0.1,
|
||||
max_tokens=1000,
|
||||
response_format={"type": "json_object"},
|
||||
api_key=api_key,
|
||||
api_base=api_base,
|
||||
)
|
||||
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
|
||||
content = result["content"]
|
||||
# Strip markdown code fences if present
|
||||
content = content.strip()
|
||||
if content.startswith("```"):
|
||||
|
||||
Reference in New Issue
Block a user