164 lines
5.4 KiB
Python
164 lines
5.4 KiB
Python
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())
|