init: 初始化项目

This commit is contained in:
zwt13703
2026-07-08 20:02:29 +08:00
parent 22590ae7b8
commit 1bb84df4ca
98 changed files with 8535 additions and 2 deletions
+181
View File
@@ -0,0 +1,181 @@
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}"
messages = [{"role": "system", "content": "你是一个专业的文档内容生成助手。"}]
messages.append({"role": "user", "content": user_content})
body = {
"model": extra.get("model", "gpt-3.5-turbo"),
"messages": messages,
"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 ("model", "messages", "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)
try:
async with httpx.AsyncClient(
timeout=timeout,
proxy=None,
trust_env=False,
) as client:
response = await client.post(
request["url"],
headers=request["headers"],
json=request["json"],
)
response.raise_for_status()
return adapter.parse_response(response.json())
except httpx.HTTPStatusError as e:
detail = e.response.text[:500] if e.response else str(e)
raise RuntimeError(f"AI 服务返回错误 ({e.response.status_code}): {detail}")
except httpx.TimeoutException:
raise RuntimeError("AI 调用超时")
except httpx.ConnectError as e:
raise RuntimeError(f"无法连接 AI 服务: {e}")
except Exception as e:
raise RuntimeError(f"AI 调用异常: {str(e)}")