125 lines
3.0 KiB
Python
125 lines
3.0 KiB
Python
from datetime import datetime
|
|
from typing import Annotated, Literal
|
|
|
|
from pydantic import BaseModel, ConfigDict, Field, field_validator, model_validator
|
|
|
|
|
|
class ToolFunctionCall(BaseModel):
|
|
name: str
|
|
arguments: str
|
|
|
|
|
|
class ToolCall(BaseModel):
|
|
id: str
|
|
type: Literal["function"] = "function"
|
|
function: ToolFunctionCall
|
|
|
|
|
|
class UserMessage(BaseModel):
|
|
role: Literal["user"] = "user"
|
|
content: str
|
|
created_at: datetime
|
|
|
|
|
|
class AssistantMessage(BaseModel):
|
|
role: Literal["assistant"] = "assistant"
|
|
content: str | None = None
|
|
tool_calls: list[ToolCall] = Field(default_factory=list)
|
|
created_at: datetime
|
|
|
|
|
|
class ToolMessage(BaseModel):
|
|
role: Literal["tool"] = "tool"
|
|
content: str
|
|
tool_call_id: str
|
|
name: str
|
|
is_error: bool = False
|
|
created_at: datetime
|
|
|
|
|
|
Message = Annotated[
|
|
UserMessage | AssistantMessage | ToolMessage,
|
|
Field(discriminator="role"),
|
|
]
|
|
class Session(BaseModel):
|
|
model_config = ConfigDict(extra="forbid")
|
|
|
|
session_id: str
|
|
system_prompt: str
|
|
created_at: datetime
|
|
updated_at: datetime
|
|
messages: list[Message] = Field(default_factory=list)
|
|
|
|
@model_validator(mode="before")
|
|
@classmethod
|
|
def add_timestamp_to_legacy_messages(cls, data: object) -> object:
|
|
if not isinstance(data, dict):
|
|
return data
|
|
|
|
fallback = data.get("created_at")
|
|
messages = data.get("messages")
|
|
if fallback is not None and isinstance(messages, list):
|
|
for message in messages:
|
|
if isinstance(message, dict):
|
|
message.setdefault("created_at", fallback)
|
|
return data
|
|
|
|
|
|
class CreateSessionRequest(BaseModel):
|
|
model_config = ConfigDict(extra="forbid")
|
|
|
|
system_prompt: str | None = None
|
|
|
|
@field_validator("system_prompt")
|
|
@classmethod
|
|
def validate_system_prompt(cls, value: str | None) -> str | None:
|
|
if value is not None and not value.strip():
|
|
raise ValueError("system_prompt must not be blank")
|
|
return value
|
|
|
|
|
|
class CreateSessionResponse(BaseModel):
|
|
session_id: str
|
|
system_prompt: str
|
|
created_at: datetime
|
|
|
|
|
|
class SendMessageRequest(BaseModel):
|
|
model_config = ConfigDict(extra="forbid")
|
|
|
|
content: str
|
|
|
|
@field_validator("content")
|
|
@classmethod
|
|
def validate_content(cls, value: str) -> str:
|
|
if not value.strip():
|
|
raise ValueError("content must not be blank")
|
|
return value
|
|
|
|
|
|
class TokenUsage(BaseModel):
|
|
prompt_tokens: int
|
|
completion_tokens: int
|
|
total_tokens: int
|
|
|
|
|
|
class FinalAssistantMessage(BaseModel):
|
|
role: Literal["assistant"] = "assistant"
|
|
content: str
|
|
created_at: datetime
|
|
|
|
|
|
class SendMessageResponse(BaseModel):
|
|
session_id: str
|
|
message: FinalAssistantMessage
|
|
tools_use: list[str] = Field(default_factory=list)
|
|
usage: TokenUsage | None
|
|
|
|
|
|
class SessionHistoryResponse(BaseModel):
|
|
session_id: str
|
|
system_prompt: str
|
|
created_at: datetime
|
|
updated_at: datetime
|
|
messages: list[Message]
|