Update sub-mode runtime model handling

This commit is contained in:
elipitc
2026-05-31 22:33:49 +08:00
parent ac7d828e80
commit 532dca13f4
7 changed files with 865 additions and 28 deletions
@@ -18,7 +18,7 @@ from fastapi.middleware.cors import CORSMiddleware
from pydantic import BaseModel, Field
import structlog
from agent import LiteLLMAgent
from agent import LiteLLMAgent, ModelRequestError
from config import get_config, AgentConfig, A2AConfig
try:
@@ -94,6 +94,7 @@ class A2ATask(BaseModel):
contextId: str = Field(default_factory=lambda: uuid.uuid4().hex)
status: A2ATaskStatus
artifacts: Optional[list[A2AArtifact]] = None
metadata: Optional[Dict[str, Any]] = None
class A2AResponse(BaseModel):
@@ -392,12 +393,12 @@ class A2AAgentServer:
request_id=task_id
) as ctx:
ctx.add_tool("a2a_chat")
response_text = await agent.chat(
response = await agent.chat_result(
message=user_text,
conversation_id=context_id
)
else:
response_text = await agent.chat(
response = await agent.chat_result(
message=user_text,
conversation_id=context_id
)
@@ -411,9 +412,17 @@ class A2AAgentServer:
task.artifacts = [
A2AArtifact(
name="response",
parts=[A2APart(kind="text", text=response_text)]
parts=[A2APart(kind="text", text=response.get("content", ""))]
)
]
task.metadata = {
"newapi_request_id": response.get("request_id"),
"response_id": response.get("response_id"),
"model": response.get("model"),
"api_format": response.get("api_format"),
"endpoint": response.get("endpoint"),
"usage": response.get("usage") or {},
}
self.tasks[task_id] = task
return JSONResponse({
@@ -425,6 +434,9 @@ class A2AAgentServer:
except Exception as e:
logger.error("处理消息失败", error=str(e))
task.status = A2ATaskStatus(state="failed", message=str(e))
error_data = {}
if isinstance(e, ModelRequestError):
error_data = e.to_dict()
self.tasks[task_id] = task
return JSONResponse({
@@ -432,7 +444,8 @@ class A2AAgentServer:
"id": request.id,
"error": {
"code": -32000,
"message": f"Agent error: {str(e)}"
"message": f"Agent error: {str(e)}",
"data": error_data,
}
})
+347 -11
View File
@@ -6,7 +6,7 @@ LiteLLM Agent 核心模块
import asyncio
import json
import uuid
from typing import AsyncGenerator, Optional, Dict, Any, List
from typing import AsyncGenerator, Optional, Dict, Any, List, Union
from dataclasses import dataclass, field
from datetime import datetime
@@ -47,6 +47,54 @@ class Conversation:
return [{"role": m.role, "content": m.content} for m in self.messages]
@dataclass
class ModelResult:
"""Normalized model response metadata for Runtime accounting."""
content: str
usage: Dict[str, int] = field(default_factory=dict)
request_id: Optional[str] = None
response_id: Optional[str] = None
model: Optional[str] = None
api_format: str = "openai_chat"
endpoint: Optional[str] = None
def to_dict(self) -> Dict[str, Any]:
return {
"content": self.content,
"usage": self.usage,
"request_id": self.request_id,
"response_id": self.response_id,
"model": self.model,
"api_format": self.api_format,
"endpoint": self.endpoint,
}
class ModelRequestError(RuntimeError):
"""Model gateway error carrying request metadata for Runtime logs."""
def __init__(
self,
message: str,
*,
request_id: Optional[str] = None,
status_code: Optional[int] = None,
response_text: Optional[str] = None,
):
super().__init__(message)
self.request_id = request_id
self.status_code = status_code
self.response_text = response_text
def to_dict(self) -> Dict[str, Any]:
return {
"request_id": self.request_id,
"status_code": self.status_code,
"response_text": self.response_text,
}
class LiteLLMAgent:
"""
基于LiteLLM的Agent实现
@@ -110,6 +158,8 @@ class LiteLLMAgent:
timeout=httpx.Timeout(self.llm_config.timeout),
headers={
"Authorization": f"Bearer {self.llm_config.api_key}",
"x-api-key": self.llm_config.api_key,
"anthropic-version": "2023-06-01",
"Content-Type": "application/json"
}
)
@@ -144,7 +194,7 @@ class LiteLLMAgent:
message: str,
conversation_id: Optional[str] = None,
stream: bool = False
) -> str | AsyncGenerator[str, None]:
) -> Union[str, AsyncGenerator[str, None]]:
"""
发送消息并获取回复
@@ -162,12 +212,73 @@ class LiteLLMAgent:
conversation.add_message("user", message)
if stream:
if self.llm_config.api_format == "anthropic_messages":
return self._stream_anthropic_messages_text(conversation)
return self._stream_chat(conversation)
else:
return await self._simple_chat(conversation)
result = await self.chat_result_for_conversation(conversation)
return result.content
async def chat_result(
self,
message: str,
conversation_id: Optional[str] = None,
) -> Dict[str, Any]:
"""Return assistant text plus usage and NewAPI request metadata."""
conversation = self.get_or_create_conversation(conversation_id)
conversation.add_message("user", message)
return (await self.chat_result_for_conversation(conversation)).to_dict()
async def chat_result_for_conversation(self, conversation: Conversation) -> ModelResult:
"""Dispatch to the configured model API format."""
if self.llm_config.api_format == "anthropic_messages":
if self.llm_config.use_stream:
return await self._anthropic_messages_stream(conversation)
return await self._anthropic_messages(conversation)
if self.llm_config.use_stream:
return await self._openai_chat_stream_result(conversation)
return await self._simple_chat(conversation)
async def _simple_chat(self, conversation: Conversation) -> str:
"""非流式对话"""
def _request_id_from_response(self, response: httpx.Response, body: Optional[Dict[str, Any]] = None) -> Optional[str]:
"""Extract NewAPI/OpenAI/Anthropic request ID from headers or body."""
for name in (
"x-request-id",
"request-id",
"x-newapi-request-id",
"x-litellm-request-id",
"anthropic-request-id",
):
value = response.headers.get(name)
if value:
return value
if body:
return body.get("request_id")
return None
def _normalize_usage(self, usage: Optional[Dict[str, Any]]) -> Dict[str, int]:
usage = usage or {}
prompt_tokens = int(usage.get("prompt_tokens") or usage.get("input_tokens") or 0)
completion_tokens = int(usage.get("completion_tokens") or usage.get("output_tokens") or 0)
if "input_tokens" in usage or "output_tokens" in usage:
total_tokens = prompt_tokens + completion_tokens
else:
total_tokens = int(usage.get("total_tokens") or prompt_tokens + completion_tokens)
return {
"prompt_tokens": prompt_tokens,
"completion_tokens": completion_tokens,
"total_tokens": total_tokens,
}
def _raise_gateway_error(self, exc: httpx.HTTPStatusError, body: Optional[Dict[str, Any]] = None) -> None:
request_id = self._request_id_from_response(exc.response, body)
raise ModelRequestError(
f"Server error '{exc.response.status_code} {exc.response.reason_phrase}' for url '{exc.request.url}'",
request_id=request_id,
status_code=exc.response.status_code,
response_text=exc.response.text[:2000],
) from exc
async def _simple_chat(self, conversation: Conversation) -> ModelResult:
"""非流式 OpenAI chat completions 对话"""
client = await self._get_client()
request_body = {
@@ -184,7 +295,10 @@ class LiteLLMAgent:
self.llm_config.chat_endpoint,
json=request_body
)
response.raise_for_status()
try:
response.raise_for_status()
except httpx.HTTPStatusError as exc:
self._raise_gateway_error(exc)
result = response.json()
assistant_message = result["choices"][0]["message"]["content"]
@@ -192,8 +306,17 @@ class LiteLLMAgent:
# 保存助手回复到对话
conversation.add_message("assistant", assistant_message)
logger.info("收到回复", length=len(assistant_message))
return assistant_message
request_id = self._request_id_from_response(response, result)
logger.info("收到回复", length=len(assistant_message), request_id=request_id)
return ModelResult(
content=assistant_message,
usage=self._normalize_usage(result.get("usage")),
request_id=request_id,
response_id=result.get("id"),
model=result.get("model") or self.llm_config.model,
api_format="openai_chat",
endpoint=self.llm_config.chat_endpoint,
)
except httpx.HTTPStatusError as e:
logger.error("HTTP错误", status_code=e.response.status_code, detail=e.response.text)
@@ -201,6 +324,215 @@ class LiteLLMAgent:
except Exception as e:
logger.error("请求失败", error=str(e))
raise
async def _openai_chat_stream_result(self, conversation: Conversation) -> ModelResult:
"""OpenAI chat completions stream=true, aggregated into one Runtime artifact."""
client = await self._get_client()
request_body = {
"model": self.llm_config.model,
"messages": conversation.to_openai_format(),
"temperature": self.llm_config.temperature,
"max_tokens": self.llm_config.max_tokens,
"stream": True,
"stream_options": {"include_usage": True},
}
full_response = ""
usage: Dict[str, int] = {}
response_id: Optional[str] = None
request_id: Optional[str] = None
try:
async with client.stream("POST", self.llm_config.chat_endpoint, json=request_body) as response:
try:
response.raise_for_status()
except httpx.HTTPStatusError as exc:
self._raise_gateway_error(exc)
request_id = self._request_id_from_response(response)
async for line in response.aiter_lines():
if not line.startswith("data: "):
continue
data = line[6:]
if data == "[DONE]":
break
try:
chunk = json.loads(data)
except json.JSONDecodeError:
continue
response_id = response_id or chunk.get("id")
request_id = request_id or chunk.get("request_id")
if chunk.get("usage"):
usage = self._normalize_usage(chunk.get("usage"))
choices = chunk.get("choices") or []
if not choices:
continue
delta = choices[0].get("delta", {})
content = delta.get("content", "")
if content:
full_response += content
conversation.add_message("assistant", full_response)
logger.info("收到流式回复", length=len(full_response), request_id=request_id)
return ModelResult(
content=full_response,
usage=usage,
request_id=request_id or response_id,
response_id=response_id,
model=self.llm_config.model,
api_format="openai_chat",
endpoint=self.llm_config.chat_endpoint,
)
except Exception as e:
logger.error("流式请求失败", error=str(e))
raise
def _anthropic_payload(self, conversation: Conversation, *, stream: bool = False) -> Dict[str, Any]:
system_parts: List[str] = []
messages: List[Dict[str, str]] = []
for message in conversation.messages:
if message.role == "system":
system_parts.append(message.content)
else:
role = "assistant" if message.role == "assistant" else "user"
messages.append({"role": role, "content": message.content})
payload: Dict[str, Any] = {
"model": self.llm_config.model,
"messages": messages,
"max_tokens": self.llm_config.max_tokens,
"stream": stream,
}
if system_parts:
payload["system"] = "\n\n".join(system_parts)
return payload
def _anthropic_text(self, body: Dict[str, Any]) -> str:
content = body.get("content") or []
texts = [
part.get("text", "")
for part in content
if isinstance(part, dict) and part.get("type") == "text"
]
return "".join(texts)
async def _anthropic_messages(self, conversation: Conversation) -> ModelResult:
"""Anthropic Messages-compatible call for Claude models."""
client = await self._get_client()
try:
response = await client.post(self.llm_config.messages_endpoint, json=self._anthropic_payload(conversation))
try:
response.raise_for_status()
except httpx.HTTPStatusError as exc:
self._raise_gateway_error(exc)
result = response.json()
assistant_message = self._anthropic_text(result)
conversation.add_message("assistant", assistant_message)
request_id = self._request_id_from_response(response, result)
logger.info("收到 Claude Messages 回复", length=len(assistant_message), request_id=request_id)
return ModelResult(
content=assistant_message,
usage=self._normalize_usage(result.get("usage")),
request_id=request_id,
response_id=result.get("id"),
model=result.get("model") or self.llm_config.model,
api_format="anthropic_messages",
endpoint=self.llm_config.messages_endpoint,
)
except Exception as e:
logger.error("Claude Messages 请求失败", error=str(e))
raise
async def _anthropic_messages_stream(self, conversation: Conversation) -> ModelResult:
"""Anthropic Messages stream=true, aggregated into one Runtime artifact."""
client = await self._get_client()
full_response = ""
usage: Dict[str, int] = {}
response_id: Optional[str] = None
request_id: Optional[str] = None
try:
async with client.stream(
"POST",
self.llm_config.messages_endpoint,
json=self._anthropic_payload(conversation, stream=True),
) as response:
try:
response.raise_for_status()
except httpx.HTTPStatusError as exc:
self._raise_gateway_error(exc)
request_id = self._request_id_from_response(response)
async for line in response.aiter_lines():
if not line.startswith("data: "):
continue
data = line[6:]
if data == "[DONE]":
break
try:
event = json.loads(data)
except json.JSONDecodeError:
continue
event_type = event.get("type")
if event_type == "message_start":
message = event.get("message") or {}
response_id = response_id or message.get("id")
usage = self._normalize_usage(message.get("usage"))
elif event_type == "content_block_delta":
delta = event.get("delta") or {}
text = delta.get("text", "")
if text:
full_response += text
elif event_type == "message_delta":
delta_usage = (event.get("usage") or {})
if delta_usage:
usage = self._normalize_usage({**usage, **delta_usage})
conversation.add_message("assistant", full_response)
logger.info("收到 Claude Messages 流式回复", length=len(full_response), request_id=request_id)
return ModelResult(
content=full_response,
usage=usage,
request_id=request_id or response_id,
response_id=response_id,
model=self.llm_config.model,
api_format="anthropic_messages",
endpoint=self.llm_config.messages_endpoint,
)
except Exception as e:
logger.error("Claude Messages 流式请求失败", error=str(e))
raise
async def _stream_anthropic_messages_text(self, conversation: Conversation) -> AsyncGenerator[str, None]:
"""Yield text deltas from Anthropic Messages stream for A2A stream clients."""
client = await self._get_client()
full_response = ""
try:
async with client.stream(
"POST",
self.llm_config.messages_endpoint,
json=self._anthropic_payload(conversation, stream=True),
) as response:
try:
response.raise_for_status()
except httpx.HTTPStatusError as exc:
self._raise_gateway_error(exc)
async for line in response.aiter_lines():
if not line.startswith("data: "):
continue
data = line[6:]
if data == "[DONE]":
break
try:
event = json.loads(data)
except json.JSONDecodeError:
continue
if event.get("type") != "content_block_delta":
continue
delta = event.get("delta") or {}
text = delta.get("text", "")
if text:
full_response += text
yield text
conversation.add_message("assistant", full_response)
except Exception as e:
logger.error("Claude Messages 文本流失败", error=str(e))
raise
async def _stream_chat(self, conversation: Conversation) -> AsyncGenerator[str, None]:
"""流式对话"""
@@ -211,7 +543,8 @@ class LiteLLMAgent:
"messages": conversation.to_openai_format(),
"temperature": self.llm_config.temperature,
"max_tokens": self.llm_config.max_tokens,
"stream": True
"stream": True,
"stream_options": {"include_usage": True},
}
full_response = ""
@@ -232,7 +565,10 @@ class LiteLLMAgent:
try:
chunk = json.loads(data)
delta = chunk.get("choices", [{}])[0].get("delta", {})
choices = chunk.get("choices") or []
if not choices:
continue
delta = choices[0].get("delta", {})
content = delta.get("content", "")
if content:
full_response += content
@@ -18,8 +18,9 @@ class LiteLLMConfig:
# 基础URL - 用户提供的LiteLLM服务地址
base_url: str = "https://litellm.graystone-fb459c5d.southeastasia.azurecontainerapps.io"
# 完整的chat completions端点
# 完整的模型端点
chat_endpoint: str = field(init=False)
messages_endpoint: str = field(init=False)
# API密钥 - 优先使用传入的,否则从环境变量获取
api_key: Optional[str] = None
@@ -38,6 +39,12 @@ class LiteLLMConfig:
# 最大token数
max_tokens: int = 4096
# API格式:openai_chat 或 anthropic_messages
api_format: str = "openai_chat"
# 是否强制使用流式请求聚合完整响应
use_stream: bool = False
def __post_init__(self):
self.base_url = (
@@ -52,7 +59,35 @@ class LiteLLMConfig:
self.api_key = os.getenv("LITELLM_API_KEY")
if self.model is None:
self.model = os.getenv("MODEL_NAME") or os.getenv("LITELLM_MODEL", "gpt-4")
model_name = (self.model or "").lower()
self.api_format = (
os.getenv("LLM_API_FORMAT")
or os.getenv("MODEL_API_FORMAT")
or ("anthropic_messages" if "claude" in model_name else "openai_chat")
).lower()
if "gpt-5.4" in model_name:
self.timeout = 600
self.use_stream = True
if "claude" in model_name:
self.timeout = 600
self.use_stream = True
if os.getenv("LITELLM_TIMEOUT") or os.getenv("LLM_TIMEOUT"):
self.timeout = int(os.getenv("LITELLM_TIMEOUT") or os.getenv("LLM_TIMEOUT"))
if os.getenv("LITELLM_MAX_TOKENS") or os.getenv("LLM_MAX_TOKENS"):
self.max_tokens = int(os.getenv("LITELLM_MAX_TOKENS") or os.getenv("LLM_MAX_TOKENS"))
if os.getenv("LITELLM_STREAM") or os.getenv("LLM_STREAM"):
self.use_stream = (os.getenv("LITELLM_STREAM") or os.getenv("LLM_STREAM", "")).lower() in {
"1",
"true",
"yes",
"on",
}
self.chat_endpoint = f"{self.base_url}/chat/completions"
self.messages_endpoint = f"{self.base_url}/messages"
def validate(self) -> bool:
"""验证配置是否完整"""