增加工具使用功能
This commit is contained in:
+141
-65
@@ -1,15 +1,24 @@
|
||||
import asyncio
|
||||
from datetime import UTC, datetime
|
||||
from typing import Any
|
||||
|
||||
from openai import APIError, AsyncOpenAI
|
||||
|
||||
from models import (
|
||||
AssistantMessage,
|
||||
FinalAssistantMessage,
|
||||
Message,
|
||||
SendMessageResponse,
|
||||
Session,
|
||||
SessionHistoryResponse,
|
||||
TokenUsage,
|
||||
ToolCall,
|
||||
ToolFunctionCall,
|
||||
ToolMessage,
|
||||
UserMessage,
|
||||
)
|
||||
from storage import JsonSessionStorage
|
||||
from tools import ToolRegistry
|
||||
|
||||
|
||||
class ChatProviderError(Exception):
|
||||
@@ -22,10 +31,18 @@ class ChatService:
|
||||
storage: JsonSessionStorage,
|
||||
client: AsyncOpenAI,
|
||||
model: str,
|
||||
tool_registry: ToolRegistry | None = None,
|
||||
max_tool_rounds: int = 5,
|
||||
max_tool_calls_per_turn: int = 10,
|
||||
) -> None:
|
||||
if max_tool_rounds <= 0 or max_tool_calls_per_turn <= 0:
|
||||
raise ValueError("tool limits must be greater than zero")
|
||||
self.storage = storage
|
||||
self.client = client
|
||||
self.model = model
|
||||
self.tool_registry = tool_registry or ToolRegistry()
|
||||
self.max_tool_rounds = max_tool_rounds
|
||||
self.max_tool_calls_per_turn = max_tool_calls_per_turn
|
||||
self._locks: dict[str, asyncio.Lock] = {}
|
||||
|
||||
async def generate_response(
|
||||
@@ -33,33 +50,88 @@ class ChatService:
|
||||
) -> SendMessageResponse:
|
||||
lock = self._locks.setdefault(session_id, asyncio.Lock())
|
||||
async with lock:
|
||||
user_created_at = datetime.now(UTC)
|
||||
user_message = UserMessage(content=content, created_at=datetime.now(UTC))
|
||||
session = await self.storage.read(session_id)
|
||||
api_messages = self._build_api_messages(session, content)
|
||||
api_messages = self._build_api_messages(session)
|
||||
api_messages.append(self._message_to_api(user_message))
|
||||
pending_messages: list[Message] = [user_message]
|
||||
tools_use: list[str] = []
|
||||
usage = TokenUsage(prompt_tokens=0, completion_tokens=0, total_tokens=0)
|
||||
has_usage = False
|
||||
tool_call_count = 0
|
||||
|
||||
try:
|
||||
completion = await self.client.chat.completions.create(
|
||||
model=self.model,
|
||||
messages=api_messages, # type: ignore[arg-type]
|
||||
extra_body={"thinking": {"type": "disabled"}},
|
||||
for _ in range(self.max_tool_rounds):
|
||||
completion = await self._request_completion(api_messages)
|
||||
if completion.usage is not None:
|
||||
self._add_usage(usage, completion.usage)
|
||||
has_usage = True
|
||||
|
||||
provider_message = completion.choices[0].message
|
||||
provider_tool_calls = getattr(provider_message, "tool_calls", None) or []
|
||||
if not provider_tool_calls:
|
||||
assistant_content = provider_message.content
|
||||
if not assistant_content:
|
||||
raise ChatProviderError(
|
||||
"Upstream chat provider returned an empty response"
|
||||
)
|
||||
final_message = AssistantMessage(
|
||||
content=assistant_content,
|
||||
created_at=datetime.now(UTC),
|
||||
)
|
||||
pending_messages.append(final_message)
|
||||
session.messages.extend(pending_messages)
|
||||
session.updated_at = final_message.created_at
|
||||
await self.storage.write(session)
|
||||
return SendMessageResponse(
|
||||
session_id=session_id,
|
||||
message=FinalAssistantMessage(
|
||||
content=assistant_content,
|
||||
created_at=final_message.created_at,
|
||||
),
|
||||
tools_use=tools_use,
|
||||
usage=usage if has_usage else None,
|
||||
)
|
||||
|
||||
tool_call_count += len(provider_tool_calls)
|
||||
if tool_call_count > self.max_tool_calls_per_turn:
|
||||
raise ChatProviderError("Tool call limit exceeded")
|
||||
|
||||
tool_calls = [
|
||||
ToolCall(
|
||||
id=tool_call.id,
|
||||
function=ToolFunctionCall(
|
||||
name=tool_call.function.name,
|
||||
arguments=tool_call.function.arguments,
|
||||
),
|
||||
)
|
||||
for tool_call in provider_tool_calls
|
||||
]
|
||||
assistant_tool_message = AssistantMessage(
|
||||
content=provider_message.content,
|
||||
tool_calls=tool_calls,
|
||||
created_at=datetime.now(UTC),
|
||||
)
|
||||
assistant_content = completion.choices[0].message.content
|
||||
except (APIError, IndexError, AttributeError) as exc:
|
||||
raise ChatProviderError("Upstream chat request failed") from exc
|
||||
pending_messages.append(assistant_tool_message)
|
||||
api_messages.append(self._message_to_api(assistant_tool_message))
|
||||
|
||||
if assistant_content is None:
|
||||
raise ChatProviderError("Upstream chat provider returned an empty response")
|
||||
for tool_call in tool_calls:
|
||||
if tool_call.function.name not in tools_use:
|
||||
tools_use.append(tool_call.function.name)
|
||||
result = await self.tool_registry.execute(
|
||||
tool_call.function.name,
|
||||
tool_call.function.arguments,
|
||||
)
|
||||
tool_message = ToolMessage(
|
||||
content=result.content,
|
||||
tool_call_id=tool_call.id,
|
||||
name=tool_call.function.name,
|
||||
is_error=result.is_error,
|
||||
created_at=datetime.now(UTC),
|
||||
)
|
||||
pending_messages.append(tool_message)
|
||||
api_messages.append(self._message_to_api(tool_message))
|
||||
|
||||
assistant_message = await self._save_exchange(
|
||||
session, content, user_created_at, assistant_content
|
||||
)
|
||||
usage = self._parse_usage(completion.usage)
|
||||
|
||||
return SendMessageResponse(
|
||||
session_id=session_id,
|
||||
message=assistant_message,
|
||||
usage=usage,
|
||||
)
|
||||
raise ChatProviderError("Tool round limit exceeded")
|
||||
|
||||
async def get_session_history(self, session_id: str) -> SessionHistoryResponse:
|
||||
lock = self._locks.setdefault(session_id, asyncio.Lock())
|
||||
@@ -73,50 +145,54 @@ class ChatService:
|
||||
messages=session.messages,
|
||||
)
|
||||
|
||||
@staticmethod
|
||||
def _build_api_messages(
|
||||
session: Session, content: str
|
||||
) -> list[dict[str, str]]:
|
||||
messages = [{"role": "system", "content": session.system_prompt}]
|
||||
messages.extend(
|
||||
{"role": message.role, "content": message.content}
|
||||
for message in session.messages
|
||||
)
|
||||
messages.append({"role": "user", "content": content})
|
||||
async def _request_completion(
|
||||
self, messages: list[dict[str, Any]]
|
||||
) -> Any:
|
||||
try:
|
||||
completion = await self.client.chat.completions.create(
|
||||
model=self.model,
|
||||
messages=messages, # type: ignore[arg-type]
|
||||
tools=self.tool_registry.definitions(), # type: ignore[arg-type]
|
||||
extra_body={"thinking": {"type": "disabled"}},
|
||||
)
|
||||
if not completion.choices:
|
||||
raise ChatProviderError("Upstream chat provider returned no choices")
|
||||
return completion
|
||||
except ChatProviderError:
|
||||
raise
|
||||
except (APIError, IndexError, AttributeError) as exc:
|
||||
raise ChatProviderError("Upstream chat request failed") from exc
|
||||
|
||||
@classmethod
|
||||
def _build_api_messages(cls, session: Session) -> list[dict[str, Any]]:
|
||||
messages: list[dict[str, Any]] = [
|
||||
{"role": "system", "content": session.system_prompt}
|
||||
]
|
||||
messages.extend(cls._message_to_api(message) for message in session.messages)
|
||||
return messages
|
||||
|
||||
async def _save_exchange(
|
||||
self,
|
||||
session: Session,
|
||||
user_content: str,
|
||||
user_created_at: datetime,
|
||||
assistant_content: str,
|
||||
) -> Message:
|
||||
assistant_message = Message(
|
||||
role="assistant",
|
||||
content=assistant_content,
|
||||
created_at=datetime.now(UTC),
|
||||
)
|
||||
session.messages.extend(
|
||||
[
|
||||
Message(
|
||||
role="user",
|
||||
content=user_content,
|
||||
created_at=user_created_at,
|
||||
),
|
||||
assistant_message,
|
||||
]
|
||||
)
|
||||
session.updated_at = assistant_message.created_at
|
||||
await self.storage.write(session)
|
||||
return assistant_message
|
||||
@staticmethod
|
||||
def _message_to_api(message: Message) -> dict[str, Any]:
|
||||
if isinstance(message, UserMessage):
|
||||
return {"role": "user", "content": message.content}
|
||||
if isinstance(message, AssistantMessage):
|
||||
result: dict[str, Any] = {
|
||||
"role": "assistant",
|
||||
"content": message.content,
|
||||
}
|
||||
if message.tool_calls:
|
||||
result["tool_calls"] = [
|
||||
tool_call.model_dump() for tool_call in message.tool_calls
|
||||
]
|
||||
return result
|
||||
return {
|
||||
"role": "tool",
|
||||
"content": message.content,
|
||||
"tool_call_id": message.tool_call_id,
|
||||
}
|
||||
|
||||
@staticmethod
|
||||
def _parse_usage(usage: object | None) -> TokenUsage | None:
|
||||
if usage is None:
|
||||
return None
|
||||
return TokenUsage(
|
||||
prompt_tokens=usage.prompt_tokens, # type: ignore[attr-defined]
|
||||
completion_tokens=usage.completion_tokens, # type: ignore[attr-defined]
|
||||
total_tokens=usage.total_tokens, # type: ignore[attr-defined]
|
||||
)
|
||||
def _add_usage(total: TokenUsage, usage: object) -> None:
|
||||
total.prompt_tokens += usage.prompt_tokens # type: ignore[attr-defined]
|
||||
total.completion_tokens += usage.completion_tokens # type: ignore[attr-defined]
|
||||
total.total_tokens += usage.total_tokens # type: ignore[attr-defined]
|
||||
|
||||
Reference in New Issue
Block a user