From d535e61f29bd32b6ad60894bf3117049c56c9ec2 Mon Sep 17 00:00:00 2001 From: yumoqing Date: Sat, 1 Aug 2026 10:25:33 +0800 Subject: [PATCH] fix: llm_bridge reads model config from llm DB table, falls back to env vars --- pipeline_service/llm_bridge.py | 86 +++++++++++++++++++++++++++------- 1 file changed, 68 insertions(+), 18 deletions(-) diff --git a/pipeline_service/llm_bridge.py b/pipeline_service/llm_bridge.py index be29177..b49e51f 100644 --- a/pipeline_service/llm_bridge.py +++ b/pipeline_service/llm_bridge.py @@ -1,8 +1,9 @@ -"""LLM bridge for pipeline handlers. +""" +LLM bridge for pipeline handlers. Provides a simple async interface for handlers to call LLM APIs. -Uses harnessed_agent's llm_chat under the hood when available, -falls back to direct HTTP calls. +Looks up model config from the llm database table first, +falls back to environment variables. """ import json @@ -11,15 +12,54 @@ import os logger = logging.getLogger("pipeline.llm_bridge") +# Cache: model_name -> {api_base, api_key, model_id} +_model_cache: dict = {} + + +async def _get_model_config(model_name: str = None) -> dict: + """Look up model config from llm table. Returns dict with api_base, api_key, model_id.""" + global _model_cache + + if model_name and model_name in _model_cache: + return _model_cache[model_name] + + try: + from sqlor.dbpools import DBPools + db = DBPools() + dbname = "pipeline" + async with db.sqlorContext(dbname) as sor: + if model_name: + sql = "SELECT api_base, api_key, model_id FROM llm WHERE name=${name}$ AND status='active' LIMIT 1" + recs = await sor.sqlExe(sql, {"name": model_name}) + else: + sql = "SELECT api_base, api_key, model_id, name FROM llm WHERE status='active' ORDER BY id LIMIT 1" + recs = await sor.sqlExe(sql, {}) + if recs: + r = recs[0] + cfg = { + "api_base": getattr(r, "api_base", "") or "", + "api_key": getattr(r, "api_key", "") or "", + "model_id": getattr(r, "model_id", "") or "", + } + cache_key = model_name or getattr(r, "name", "") + if cache_key: + _model_cache[cache_key] = cfg + return cfg + except Exception as e: + logger.warning("llm_bridge: DB lookup failed: %s", e) + + return {} + async def llm_call(prompt: str, model: str = None, temperature: float = 0.7) -> str: """Call LLM and return text response. - - Tries multiple backends: + + Backend priority: 1. harnessed_agent.llm_chat (if loaded in ServerEnv) - 2. Direct OpenAI-compatible API call + 2. DB llm table (api_base + api_key) + 3. Environment variables (LLM_API_BASE, LLM_API_KEY, LLM_MODEL) """ - # Try harnessed_agent first + # Priority 1: harnessed_agent try: from ahserver.serverenv import ServerEnv env = ServerEnv() @@ -31,32 +71,42 @@ async def llm_call(prompt: str, model: str = None, temperature: float = 0.7) -> except Exception: pass - # Fallback: direct HTTP call to OpenAI-compatible endpoint - import aiohttp - - api_base = os.environ.get("LLM_API_BASE", "https://api.openai.com/v1") - api_key = os.environ.get("LLM_API_KEY", "") - model = model or os.environ.get("LLM_MODEL", "gpt-4o-mini") - + # Priority 2: DB llm table + cfg = await _get_model_config(model) + if cfg.get("api_key") and cfg.get("api_base"): + api_base = cfg["api_base"] + api_key = cfg["api_key"] + model_id = cfg.get("model_id") or model or "default" + logger.info("llm_bridge: using DB model config for %s -> %s", model, api_base) + else: + # Priority 3: Environment variables + api_base = os.environ.get("LLM_API_BASE", "https://api.openai.com/v1") + api_key = os.environ.get("LLM_API_KEY", "") + model_id = model or os.environ.get("LLM_MODEL", "gpt-4o-mini") + if not api_key: - raise ValueError("No LLM API configured (set LLM_API_KEY env var)") + raise ValueError("No LLM API configured. Please add a model in the llm table or set LLM_API_KEY env var.") + + import aiohttp headers = { "Authorization": f"Bearer {api_key}", "Content-Type": "application/json", } payload = { - "model": model, + "model": model_id, "messages": [{"role": "user", "content": prompt}], "temperature": temperature, } + url = api_base.rstrip("/") + "/chat/completions" async with aiohttp.ClientSession() as session: async with session.post( - f"{api_base}/chat/completions", headers=headers, json=payload, timeout=aiohttp.ClientTimeout(total=120) + url, headers=headers, json=payload, + timeout=aiohttp.ClientTimeout(total=120) ) as resp: if resp.status != 200: text = await resp.text() - raise ValueError(f"LLM API error {resp.status}: {text[:200]}") + raise ValueError(f"LLM API error {resp.status}: {text[:300]}") data = await resp.json() return data["choices"][0]["message"]["content"]