first commit
This commit is contained in:
@@ -0,0 +1,327 @@
|
||||
"""
|
||||
Chat router module.
|
||||
"""
|
||||
|
||||
import time
|
||||
import json
|
||||
from typing import AsyncGenerator, Any
|
||||
from SimpleLLMFunc import push_error
|
||||
from fastapi import APIRouter, Request, Depends, HTTPException
|
||||
from fastapi.responses import StreamingResponse, JSONResponse
|
||||
|
||||
from ..models import (
|
||||
ChatCompletionRequest,
|
||||
)
|
||||
from context.conversation_manager import Conversation
|
||||
from ..state import get_server_state, ServerState
|
||||
from ..utils import (
|
||||
validate_chat_request,
|
||||
get_or_create_conversation,
|
||||
get_agent_for_model,
|
||||
process_agent_response,
|
||||
create_chat_response,
|
||||
persist_request_images,
|
||||
)
|
||||
from ..error_handlers import create_error_response
|
||||
from SimpleLLMFunc.logger import (
|
||||
app_log,
|
||||
push_warning,
|
||||
log_context,
|
||||
get_location,
|
||||
)
|
||||
from agent import BaseAgent
|
||||
from observability import propagate_conversation_session
|
||||
from react_stream import (
|
||||
event_name_for_output,
|
||||
format_sse,
|
||||
is_response_yield,
|
||||
project_response_to_oai_chunk,
|
||||
serialize_react_output,
|
||||
)
|
||||
|
||||
router = APIRouter(prefix="/v1/chat", tags=["chat"])
|
||||
|
||||
|
||||
async def _persist_conversation(conversation: Conversation) -> None:
|
||||
try:
|
||||
await conversation.context.persist()
|
||||
conversation.sketch_pad.persist()
|
||||
app_log(
|
||||
f"✅ Auto-saved conversation {conversation.uuid} after stream completion"
|
||||
)
|
||||
except Exception as save_error:
|
||||
push_warning(
|
||||
f"⚠️ Warning: Failed to save conversation {conversation.uuid}: {save_error}"
|
||||
)
|
||||
|
||||
|
||||
async def stream_chat_completion(
|
||||
request: ChatCompletionRequest,
|
||||
request_id: str,
|
||||
conversation: Conversation,
|
||||
agent: BaseAgent,
|
||||
) -> AsyncGenerator[str, None]:
|
||||
"""OpenAI-compatible streaming projection built from ReactOutput."""
|
||||
query, _, raw_user_content = validate_chat_request(request)
|
||||
|
||||
app_log(
|
||||
f"🔍 Starting stream chat completion for conversation {conversation.uuid}, the query is {query}, agent is {agent.name}"
|
||||
)
|
||||
|
||||
created_time = int(time.time())
|
||||
try:
|
||||
with propagate_conversation_session(
|
||||
conversation_id=conversation.uuid,
|
||||
metadata={
|
||||
"model": request.model,
|
||||
"agent_name": agent.name,
|
||||
"request_id": request_id,
|
||||
"transport": "oai_stream",
|
||||
},
|
||||
tags=["cadagent", "oai_stream"],
|
||||
):
|
||||
with conversation:
|
||||
sent_role = False
|
||||
async for output in agent.run(query, raw_user_content=raw_user_content):
|
||||
if not is_response_yield(output):
|
||||
continue
|
||||
|
||||
chunk_obj = project_response_to_oai_chunk(
|
||||
output,
|
||||
request_id=request_id,
|
||||
model=request.model,
|
||||
created_time=created_time,
|
||||
sent_role=sent_role,
|
||||
)
|
||||
if chunk_obj is None:
|
||||
continue
|
||||
|
||||
try:
|
||||
json_str = json.dumps(chunk_obj, ensure_ascii=False)
|
||||
except Exception as encode_err:
|
||||
push_warning(f"Failed to encode projected chunk: {encode_err}")
|
||||
json_str = json.dumps(
|
||||
{"error": str(encode_err)}, ensure_ascii=False
|
||||
)
|
||||
|
||||
first_delta = chunk_obj.get("choices", [{}])[0].get("delta", {})
|
||||
if (
|
||||
isinstance(first_delta, dict)
|
||||
and first_delta.get("role") == "assistant"
|
||||
):
|
||||
sent_role = True
|
||||
|
||||
app_log(f"🔍 Forwarding projected chunk: {json_str}")
|
||||
yield f"data: {json_str}\n\n"
|
||||
|
||||
await _persist_conversation(conversation)
|
||||
|
||||
except Exception as e:
|
||||
err_obj: dict[str, Any] = {
|
||||
"id": request_id,
|
||||
"object": "chat.completion.chunk",
|
||||
"created": created_time,
|
||||
"model": request.model,
|
||||
"choices": [
|
||||
{
|
||||
"index": 0,
|
||||
"delta": {"role": "assistant", "content": f"Error: {str(e)}"},
|
||||
"finish_reason": "stop",
|
||||
}
|
||||
],
|
||||
}
|
||||
err_json = json.dumps(err_obj, ensure_ascii=False)
|
||||
push_error(f"🔍 Sending error chunk: {err_json}", location=get_location())
|
||||
yield f"data: {err_json}\n\n"
|
||||
yield "data: [DONE]\n\n"
|
||||
return
|
||||
|
||||
# End signal.
|
||||
yield "data: [DONE]\n\n"
|
||||
|
||||
|
||||
async def stream_chat_events(
|
||||
request: ChatCompletionRequest,
|
||||
conversation: Conversation,
|
||||
agent: BaseAgent,
|
||||
) -> AsyncGenerator[str, None]:
|
||||
"""Native SSE stream that exposes our React event protocol."""
|
||||
query, request_id, raw_user_content = validate_chat_request(request)
|
||||
|
||||
try:
|
||||
with propagate_conversation_session(
|
||||
conversation_id=conversation.uuid,
|
||||
metadata={
|
||||
"model": request.model,
|
||||
"agent_name": agent.name,
|
||||
"request_id": request_id,
|
||||
"transport": "event_stream",
|
||||
},
|
||||
tags=["cadagent", "event_stream"],
|
||||
):
|
||||
with conversation:
|
||||
async for output in agent.run(query, raw_user_content=raw_user_content):
|
||||
payload = serialize_react_output(output, delta_consumer="web")
|
||||
yield format_sse(event_name_for_output(output), payload)
|
||||
|
||||
await _persist_conversation(conversation)
|
||||
except Exception as e:
|
||||
error_payload = {
|
||||
"type": "error",
|
||||
"message": str(e),
|
||||
}
|
||||
yield format_sse("error", error_payload)
|
||||
yield format_sse("done", {"ok": False})
|
||||
return
|
||||
|
||||
yield format_sse("done", {"ok": True})
|
||||
|
||||
|
||||
@router.post(
|
||||
"/completions",
|
||||
dependencies=[Depends(get_server_state)],
|
||||
description="Chat completion endpoint compatible with the OpenAI specification",
|
||||
)
|
||||
async def chat_completions(
|
||||
request: ChatCompletionRequest,
|
||||
http_request: Request,
|
||||
state: ServerState = Depends(get_server_state),
|
||||
):
|
||||
"""Chat completion endpoint compatible with the OpenAI specification."""
|
||||
if not state.agent_registry:
|
||||
return create_error_response(
|
||||
message="Agent registry not initialized",
|
||||
error_type="server_error",
|
||||
status_code=500,
|
||||
)
|
||||
|
||||
if not state.conversation_manager:
|
||||
return create_error_response(
|
||||
message="Conversation manager not initialized",
|
||||
error_type="server_error",
|
||||
status_code=500,
|
||||
)
|
||||
|
||||
# Get conversation ID from the custom header.
|
||||
conversation_id = http_request.headers.get("X-Conversation-ID")
|
||||
|
||||
with log_context(conversation_id=conversation_id):
|
||||
try:
|
||||
# Get or create the conversation.
|
||||
conversation, conversation_id = get_or_create_conversation(
|
||||
conversation_id, state.conversation_manager
|
||||
)
|
||||
|
||||
persist_request_images(request, conversation_id)
|
||||
|
||||
# Get Agent.
|
||||
agent = get_agent_for_model(request.model, state.agent_registry)
|
||||
|
||||
# Validate the request and retrieve the necessary information.
|
||||
query, request_id, raw_user_content = validate_chat_request(request)
|
||||
|
||||
app_log(
|
||||
f"🔍 {request_id} request chat completion for conversation {conversation_id}"
|
||||
)
|
||||
|
||||
# Streaming response.
|
||||
if request.stream:
|
||||
return StreamingResponse(
|
||||
stream_chat_completion(request, request_id, conversation, agent),
|
||||
media_type="text/event-stream",
|
||||
headers={
|
||||
"Cache-Control": "no-cache",
|
||||
"Connection": "keep-alive",
|
||||
"Access-Control-Allow-Origin": "*",
|
||||
"X-Conversation-ID": conversation_id,
|
||||
},
|
||||
)
|
||||
|
||||
# Non-streaming response: return text plus token statistics.
|
||||
(
|
||||
full_response,
|
||||
prompt_tokens,
|
||||
completion_tokens,
|
||||
) = await process_agent_response(
|
||||
query,
|
||||
conversation,
|
||||
agent,
|
||||
raw_user_content=raw_user_content,
|
||||
)
|
||||
response = create_chat_response(
|
||||
request_id,
|
||||
request.model,
|
||||
full_response,
|
||||
prompt_tokens,
|
||||
completion_tokens,
|
||||
)
|
||||
|
||||
return JSONResponse(
|
||||
content=response.model_dump(),
|
||||
headers={"X-Conversation-ID": conversation_id},
|
||||
)
|
||||
|
||||
except HTTPException:
|
||||
# Re-raise HTTPException so FastAPI can handle it.
|
||||
raise
|
||||
except Exception as e:
|
||||
return create_error_response(
|
||||
message=f"Internal server error: {str(e)}",
|
||||
error_type="server_error",
|
||||
status_code=500,
|
||||
)
|
||||
|
||||
|
||||
@router.post(
|
||||
"/events",
|
||||
dependencies=[Depends(get_server_state)],
|
||||
description="Chat event stream endpoint using native ReactEvent Stream SSE",
|
||||
)
|
||||
async def chat_events(
|
||||
request: ChatCompletionRequest,
|
||||
http_request: Request,
|
||||
state: ServerState = Depends(get_server_state),
|
||||
):
|
||||
if not state.agent_registry:
|
||||
return create_error_response(
|
||||
message="Agent registry not initialized",
|
||||
error_type="server_error",
|
||||
status_code=500,
|
||||
)
|
||||
|
||||
if not state.conversation_manager:
|
||||
return create_error_response(
|
||||
message="Conversation manager not initialized",
|
||||
error_type="server_error",
|
||||
status_code=500,
|
||||
)
|
||||
|
||||
conversation_id = http_request.headers.get("X-Conversation-ID")
|
||||
|
||||
with log_context(conversation_id=conversation_id):
|
||||
try:
|
||||
conversation, conversation_id = get_or_create_conversation(
|
||||
conversation_id, state.conversation_manager
|
||||
)
|
||||
persist_request_images(request, conversation_id)
|
||||
agent = get_agent_for_model(request.model, state.agent_registry)
|
||||
validate_chat_request(request)
|
||||
|
||||
return StreamingResponse(
|
||||
stream_chat_events(request, conversation, agent),
|
||||
media_type="text/event-stream",
|
||||
headers={
|
||||
"Cache-Control": "no-cache",
|
||||
"Connection": "keep-alive",
|
||||
"Access-Control-Allow-Origin": "*",
|
||||
"X-Conversation-ID": conversation_id,
|
||||
},
|
||||
)
|
||||
except HTTPException:
|
||||
raise
|
||||
except Exception as e:
|
||||
return create_error_response(
|
||||
message=f"Internal server error: {str(e)}",
|
||||
error_type="server_error",
|
||||
status_code=500,
|
||||
)
|
||||
Reference in New Issue
Block a user