fix: unified_search + ai_proactive get API key from DB, fix model names for Ollama Cloud
- query_understanding.py: get API key/base_url/provider_type from ai_providers DB - embedding.py: get API key from DB, pass db+tenant_id through call chain - routes.py: pass db+tenant_id to llm_analyze_query and llm_aggregate_results - search_engine.py: pass db+tenant_id to generate_embedding - unified_search/jobs.py: pass db+tenant_id to generate_embedding - Fix all default model names: ollama/deepseek-v4 -> ollama/deepseek-v4-flash - Ollama Cloud has no embedding endpoint; embedding calls fail gracefully
This commit is contained in:
@@ -5,7 +5,7 @@ from __future__ import annotations
|
||||
import os
|
||||
import logging
|
||||
import uuid
|
||||
from typing import TYPE_CHECKING
|
||||
from typing import Any, TYPE_CHECKING
|
||||
|
||||
import litellm
|
||||
|
||||
@@ -17,15 +17,49 @@ logger = logging.getLogger(__name__)
|
||||
MAX_INPUT_CHARS = 8000
|
||||
|
||||
DEFAULT_EMBEDDING_MODEL = os.environ.get('SEARCH_EMBEDDING_MODEL', 'ollama/nomic-embed-text')
|
||||
OLLAMA_API_KEY = os.environ.get('API_KEY_OLLAMA_CLOUD', '')
|
||||
|
||||
|
||||
async def generate_embedding(text: str, model: str | None = None) -> list[float]:
|
||||
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
|
||||
|
||||
|
||||
async def generate_embedding(
|
||||
text: str,
|
||||
model: str | None = None,
|
||||
db: "AsyncSession | None" = None,
|
||||
tenant_id: "uuid.UUID | None" = None,
|
||||
) -> list[float]:
|
||||
"""Generate a single embedding via LiteLLM.
|
||||
|
||||
Args:
|
||||
text: Input text (truncated to 8000 chars).
|
||||
model: Embedding model name (default: ollama/nomic-embed-text).
|
||||
db: Optional DB session for API key lookup.
|
||||
tenant_id: Optional tenant ID for API key lookup.
|
||||
|
||||
Returns:
|
||||
Embedding vector as list of floats.
|
||||
@@ -33,25 +67,38 @@ async def generate_embedding(text: str, model: str | None = None) -> list[float]
|
||||
model = model or DEFAULT_EMBEDDING_MODEL
|
||||
truncated = text[:MAX_INPUT_CHARS]
|
||||
try:
|
||||
response = await litellm.aembedding(
|
||||
model=model,
|
||||
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,
|
||||
api_key=OLLAMA_API_KEY,
|
||||
)
|
||||
if api_key:
|
||||
litellm_kwargs["api_key"] = api_key
|
||||
if api_base:
|
||||
litellm_kwargs["api_base"] = api_base
|
||||
|
||||
response = await litellm.aembedding(**litellm_kwargs)
|
||||
return response.data[0]["embedding"]
|
||||
except Exception:
|
||||
logger.exception("Failed to generate embedding")
|
||||
logger.warning("Failed to generate embedding (Ollama Cloud may not support embeddings)", exc_info=True)
|
||||
return []
|
||||
|
||||
|
||||
async def generate_embeddings_batch(
|
||||
texts: list[str], model: str | None = None
|
||||
texts: list[str],
|
||||
model: str | None = None,
|
||||
db: "AsyncSession | None" = None,
|
||||
tenant_id: "uuid.UUID | None" = None,
|
||||
) -> list[list[float]]:
|
||||
"""Generate embeddings for multiple texts in a single API call.
|
||||
|
||||
Args:
|
||||
texts: List of input texts.
|
||||
model: Embedding model name (default: ollama/nomic-embed-text).
|
||||
db: Optional DB session for API key lookup.
|
||||
tenant_id: Optional tenant ID for API key lookup.
|
||||
|
||||
Returns:
|
||||
List of embedding vectors.
|
||||
@@ -59,14 +106,22 @@ async def generate_embeddings_batch(
|
||||
model = model or DEFAULT_EMBEDDING_MODEL
|
||||
truncated = [t[:MAX_INPUT_CHARS] for t in texts]
|
||||
try:
|
||||
response = await litellm.aembedding(
|
||||
model=model,
|
||||
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,
|
||||
api_key=OLLAMA_API_KEY,
|
||||
)
|
||||
if api_key:
|
||||
litellm_kwargs["api_key"] = api_key
|
||||
if api_base:
|
||||
litellm_kwargs["api_base"] = api_base
|
||||
|
||||
response = await litellm.aembedding(**litellm_kwargs)
|
||||
return [d["embedding"] for d in response.data]
|
||||
except Exception:
|
||||
logger.exception("Failed to generate batch embeddings")
|
||||
logger.warning("Failed to generate batch embeddings", exc_info=True)
|
||||
return [[] for _ in texts]
|
||||
|
||||
|
||||
@@ -74,7 +129,7 @@ async def index_entity(
|
||||
entity_type: str,
|
||||
entity_id: uuid.UUID,
|
||||
tenant_id: uuid.UUID,
|
||||
db: AsyncSession,
|
||||
db: "AsyncSession",
|
||||
) -> bool:
|
||||
"""Generate and store embedding for a single entity.
|
||||
|
||||
@@ -97,7 +152,7 @@ async def index_entity(
|
||||
logger.debug("Empty embedding text for %s/%s", entity_type, entity_id)
|
||||
return False
|
||||
|
||||
embedding = await generate_embedding(text)
|
||||
embedding = await generate_embedding(text, db=db, tenant_id=tenant_id)
|
||||
if not embedding:
|
||||
return False
|
||||
|
||||
|
||||
Reference in New Issue
Block a user