agent开发-日志agent开发日志 day3

接下来这段代码主要是一个llm工厂,它的主要功能有选择llm,读取环境变量+校验,读取apikey,连接池管理,还有结构化输出这几个功能

点击查看代码
"""
LLM 工厂 — 统一模型入口
========================
职责:
  1. 根据环境变量自动选择 OpenAI / DeepSeek / Ollama
  2. 统一返回 langchain ChatOpenAI 实例
  3. 集中管理 api_key / base_url / model / temperature 等参数

使用方式:
  from app.llm_config import get_llm
  llm = get_llm()
  response = llm.invoke("你好")

环境变量(放入 .env 文件或系统环境变量):
  LLM_PROVIDER   = openai | deepseek | ollama (默认 openai)
  OPENAI_API_KEY = sk-xxx
  DEEPSEEK_API_KEY = sk-xxx
  LLM_MODEL_NAME = gpt-4o-mini(覆盖默认模型名)
  LLM_BASE_URL   = 自定义 base_url(覆盖默认)
  LLM_TEMPERATURE = 0.3(默认 0.3)
  LLM_MAX_TOKENS  = 4096(默认 4096)
"""

import os
from functools import lru_cache
from typing import Literal

from dotenv import load_dotenv
from langchain_openai import ChatOpenAI

load_dotenv()

Provider = Literal["openai", "deepseek", "ollama"]

DEFAULT_CONFIGS = {
    "openai": {
        "model": "gpt-4o-mini",
        "base_url": "https://api.openai.com/v1",
        "api_key_env": "OPENAI_API_KEY",
    },
    "deepseek": {
        "model": "deepseek-v4-pro",
        "base_url": "https://api.deepseek.com",
        "api_key_env": "DEEPSEEK_API_KEY",
    },
    "ollama": {
        "model": "qwen2.5:7b",
        "base_url": "http://localhost:11434/v1",
        "api_key_env": None,
    },
}


def _get_provider() -> Provider:
    provider = os.getenv("LLM_PROVIDER", "openai").lower()
    if provider not in DEFAULT_CONFIGS:
        raise ValueError(
            f"不支持的 LLM_PROVIDER: {provider},可选: {list(DEFAULT_CONFIGS.keys())}"
        )
    return provider


def _resolve_api_key(config: dict) -> str:
    env_var = config["api_key_env"]
    if env_var is None:
        return "ollama"
    key = os.getenv(env_var, "")
    if not key:
        raise ValueError(
            f"缺少 API Key:请设置环境变量 {env_var},"
            f"或在 .env 文件中添加 {env_var}=your-key"
        )
    return key


@lru_cache()
def get_llm() -> ChatOpenAI:
    provider = _get_provider()
    config = DEFAULT_CONFIGS[provider]

    model = os.getenv("LLM_MODEL_NAME", config["model"])
    base_url = os.getenv("LLM_BASE_URL", config["base_url"])
    api_key = _resolve_api_key(config)
    temperature = float(os.getenv("LLM_TEMPERATURE", "0.3"))
    max_tokens = int(os.getenv("LLM_MAX_TOKENS", "4096"))

    return ChatOpenAI(
        model=model,
        base_url=base_url,
        api_key=api_key,
        temperature=temperature,
        max_tokens=max_tokens,
    )


def get_llm_info() -> dict:
    provider = _get_provider()
    config = DEFAULT_CONFIGS[provider]
    return {
        "provider": provider,
        "model": os.getenv("LLM_MODEL_NAME", config["model"]),
        "base_url": os.getenv("LLM_BASE_URL", config["base_url"]),
        "temperature": float(os.getenv("LLM_TEMPERATURE", "0.3")),
        "max_tokens": int(os.getenv("LLM_MAX_TOKENS", "4096")),
    }


def get_structured_llm(output_schema):
    """
    返回绑定了 structured_output 的 LLM 实例。
    
    使用 method="function_calling" 而非默认的 "json_schema":
      - DeepSeek / 部分模型不支持 response_format(json_schema 模式)
      - function_calling(工具调用模式)兼容性更好,DeepSeek 原生支持
      - LangChain 自动把 Pydantic Schema 转为 function definition 传给模型
    """
    return get_llm().with_structured_output(output_schema, method="function_calling")


__all__ = ["get_llm", "get_structured_llm", "get_llm_info", "Provider"]

posted @ 2026-05-21 11:16  marisa3  阅读(18)  评论(0)    收藏  举报