import json from abc import ABC, abstractmethod import httpx from app.core.security import decrypt_api_key PROVIDER_OPENAI = "openai" PROVIDER_AZURE = "azure" PROVIDER_CUSTOM = "custom" class AIAdapter(ABC): @abstractmethod def build_request(self, model_config: dict, prompt: str, ref_content: str | None = None) -> dict: pass @abstractmethod def parse_response(self, response_data: dict) -> str: pass @property @abstractmethod def provider(self) -> str: pass class OpenAIAdapter(AIAdapter): @property def provider(self) -> str: return PROVIDER_OPENAI def build_request(self, model_config: dict, prompt: str, ref_content: str | None = None) -> dict: extra = model_config.get("extra_params", {}) temperature = extra.get("temperature", 0.7) max_tokens = extra.get("max_tokens", 2000) messages = [{"role": "system", "content": "你是一个专业的文档内容生成助手。"}] user_content = prompt if ref_content: user_content = f"参考以下内容:\n{ref_content}\n\n任务:{prompt}" messages.append({"role": "user", "content": user_content}) return { "url": model_config["endpoint"], "headers": { "Authorization": f"Bearer {decrypt_api_key(model_config['api_key'])}", "Content-Type": "application/json", }, "json": { "model": extra.get("model", "gpt-4"), "messages": messages, "temperature": temperature, "max_tokens": max_tokens, }, } def parse_response(self, response_data: dict) -> str: return response_data["choices"][0]["message"]["content"] class AzureAdapter(AIAdapter): @property def provider(self) -> str: return PROVIDER_AZURE def build_request(self, model_config: dict, prompt: str, ref_content: str | None = None) -> dict: extra = model_config.get("extra_params", {}) temperature = extra.get("temperature", 0.7) max_tokens = extra.get("max_tokens", 2000) messages = [{"role": "system", "content": "你是一个专业的文档内容生成助手。"}] user_content = prompt if ref_content: user_content = f"参考以下内容:\n{ref_content}\n\n任务:{prompt}" messages.append({"role": "user", "content": user_content}) api_version = extra.get("api_version", "2024-02-15-preview") endpoint = model_config["endpoint"] url = f"{endpoint}?api-version={api_version}" return { "url": url, "headers": { "api-key": decrypt_api_key(model_config["api_key"]), "Content-Type": "application/json", }, "json": { "messages": messages, "temperature": temperature, "max_tokens": max_tokens, }, } def parse_response(self, response_data: dict) -> str: return response_data["choices"][0]["message"]["content"] class CustomAdapter(AIAdapter): @property def provider(self) -> str: return PROVIDER_CUSTOM def build_request(self, model_config: dict, prompt: str, ref_content: str | None = None) -> dict: extra = model_config.get("extra_params", {}) user_content = prompt if ref_content: user_content = f"参考以下内容:\n{ref_content}\n\n任务:{prompt}" body = { "prompt": user_content, "max_tokens": extra.get("max_tokens", 2000), "temperature": extra.get("temperature", 0.7), } body.update({k: v for k, v in extra.items() if k not in ("max_tokens", "temperature")}) return { "url": model_config["endpoint"], "headers": { "Authorization": f"Bearer {decrypt_api_key(model_config['api_key'])}", "Content-Type": "application/json", }, "json": body, } def parse_response(self, response_data: dict) -> str: if "choices" in response_data: return response_data["choices"][0]["message"]["content"] if "response" in response_data: return response_data["response"] if "content" in response_data: return response_data["content"] if "text" in response_data: return response_data["text"] return json.dumps(response_data) _adapters: dict[str, AIAdapter] = { PROVIDER_OPENAI: OpenAIAdapter(), PROVIDER_AZURE: AzureAdapter(), PROVIDER_CUSTOM: CustomAdapter(), } def get_adapter(provider: str) -> AIAdapter: adapter = _adapters.get(provider) if not adapter: raise ValueError(f"不支持的供应商: {provider}") return adapter async def call_ai_model(model_config: dict, prompt: str, ref_content: str | None = None) -> str: adapter = get_adapter(model_config["provider"]) request = adapter.build_request(model_config, prompt, ref_content) timeout = model_config.get("extra_params", {}).get("timeout", 120) async with httpx.AsyncClient(timeout=timeout) as client: response = await client.post( request["url"], headers=request["headers"], json=request["json"], ) response.raise_for_status() return adapter.parse_response(response.json())