chore: add system instruction to enhance compliance with function call

This commit is contained in:
Toddy
2025-02-27 10:35:25 +00:00
parent a592269198
commit 8d48db026c
2 changed files with 36 additions and 17 deletions
+21 -12
View File
@@ -1,14 +1,16 @@
# app/services/chat/message_converter.py # app/services/chat/message_converter.py
from abc import ABC, abstractmethod from abc import ABC, abstractmethod
from typing import List, Dict, Any from typing import Any, Dict, List, Optional
SUPPORTED_ROLES = ["user", "model", "system"]
class MessageConverter(ABC): class MessageConverter(ABC):
"""消息转换器基类""" """消息转换器基类"""
@abstractmethod @abstractmethod
def convert(self, messages: List[Dict[str, Any]]) -> List[Dict[str, Any]]: def convert(self, messages: List[Dict[str, Any]]) -> tuple[List[Dict[str, Any]], Optional[Dict[str, Any]]]:
pass pass
@@ -30,16 +32,19 @@ def _convert_image(image_url: str) -> Dict[str, Any]:
class OpenAIMessageConverter(MessageConverter): class OpenAIMessageConverter(MessageConverter):
"""OpenAI消息格式转换器""" """OpenAI消息格式转换器"""
def convert(self, messages: List[Dict[str, Any]]) -> List[Dict[str, Any]]: def convert(self, messages: List[Dict[str, Any]]) -> tuple[List[Dict[str, Any]], Optional[Dict[str, Any]]]:
converted_messages = [] converted_messages = []
for msg in messages: system_instruction = None
role = "user" if msg["role"] == "user" else "model"
parts = []
if isinstance(msg["content"], str): for msg in messages:
# 请求 gemini 接口时如果包含 content 字段但内容为空时会返回 400 错误,所以需要判断是否为空 role = msg.get("role", "")
if msg["content"]: if role not in SUPPORTED_ROLES:
parts.append({"text": msg["content"]}) role = "model"
parts = []
if isinstance(msg["content"], str) and msg["content"]:
# 请求 gemini 接口时如果包含 content 字段但内容为空时会返回 400 错误,所以需要判断是否为空并移除
parts.append({"text": msg["content"]})
elif isinstance(msg["content"], list): elif isinstance(msg["content"], list):
for content in msg["content"]: for content in msg["content"]:
if isinstance(content, str) and content: if isinstance(content, str) and content:
@@ -50,6 +55,10 @@ class OpenAIMessageConverter(MessageConverter):
elif content["type"] == "image_url": elif content["type"] == "image_url":
parts.append(_convert_image(content["image_url"]["url"])) parts.append(_convert_image(content["image_url"]["url"]))
converted_messages.append({"role": role, "parts": parts}) if parts:
if role == "system":
system_instruction = {"role": "system", "parts": parts}
else:
converted_messages.append({"role": role, "parts": parts})
return converted_messages return converted_messages, system_instruction
+15 -5
View File
@@ -2,7 +2,7 @@
from copy import deepcopy from copy import deepcopy
import json import json
from typing import Dict, Any, AsyncGenerator, List, Union from typing import Dict, Any, AsyncGenerator, List, Optional, Union
from app.core.logger import get_openai_logger from app.core.logger import get_openai_logger
from app.services.chat.message_converter import OpenAIMessageConverter from app.services.chat.message_converter import OpenAIMessageConverter
from app.services.chat.response_handler import OpenAIResponseHandler from app.services.chat.response_handler import OpenAIResponseHandler
@@ -87,10 +87,10 @@ def _get_safety_settings(model: str) -> List[Dict[str, str]]:
def _build_payload( def _build_payload(
request: ChatRequest, messages: List[Dict[str, Any]] request: ChatRequest, messages: List[Dict[str, Any]], instruction: Optional[Dict[str, Any]] = None
) -> Dict[str, Any]: ) -> Dict[str, Any]:
"""构建请求payload""" """构建请求payload"""
return { payload = {
"contents": messages, "contents": messages,
"generationConfig": { "generationConfig": {
"temperature": request.temperature, "temperature": request.temperature,
@@ -103,6 +103,16 @@ def _build_payload(
"safetySettings": _get_safety_settings(request.model), "safetySettings": _get_safety_settings(request.model),
} }
if (
instruction
and isinstance(instruction, dict)
and instruction.get("role") == "system"
and instruction.get("parts")
):
payload["systemInstruction"] = instruction
return payload
class OpenAIChatService: class OpenAIChatService:
"""聊天服务""" """聊天服务"""
@@ -120,10 +130,10 @@ class OpenAIChatService:
) -> Union[Dict[str, Any], AsyncGenerator[str, None]]: ) -> Union[Dict[str, Any], AsyncGenerator[str, None]]:
"""创建聊天完成""" """创建聊天完成"""
# 转换消息格式 # 转换消息格式
messages = self.message_converter.convert(request.messages) messages, instruction = self.message_converter.convert(request.messages)
# 构建请求payload # 构建请求payload
payload = _build_payload(request, messages) payload = _build_payload(request, messages, instruction)
if request.stream: if request.stream:
return self._handle_stream_completion(request.model, payload, api_key) return self._handle_stream_completion(request.model, payload, api_key)