Update sub-mode runtime model handling
This commit is contained in:
@@ -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,
|
||||
}
|
||||
})
|
||||
|
||||
|
||||
@@ -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:
|
||||
"""验证配置是否完整"""
|
||||
|
||||
Reference in New Issue
Block a user