feat: compatible api

This commit is contained in:
Alex committed 2026-03-31 23:10:09 +01:00
1 parent 73256389cf
commit addf57cab7
15 files changed
+1988 -8

No files matched your search

+3
View File
@@ -0,0 +1,3 @@
from application.api.v1.routes import v1_bp
__all__ = ["v1_bp"]
+311
View File
@@ -0,0 +1,311 @@
"""Standard chat completions API routes.
Exposes ``/v1/chat/completions`` and ``/v1/models`` endpoints that
follow the widely-adopted chat completions protocol so external tools
(opencode, continue, etc.) can connect to DocsGPT agents.
"""
import json
import logging
import time
import traceback
from typing import Any, Dict, Generator, Optional
from flask import Blueprint, jsonify, make_response, request, Response
from application.api.answer.routes.base import BaseAnswerResource
from application.api.answer.services.stream_processor import StreamProcessor
from application.api.v1.translator import (
translate_request,
translate_response,
translate_stream_event,
)
from application.core.mongo_db import MongoDB
from application.core.settings import settings
logger = logging.getLogger(__name__)
v1_bp = Blueprint("v1", __name__, url_prefix="/v1")
def _extract_bearer_token() -> Optional[str]:
"""Extract API key from Authorization: Bearer header."""
auth = request.headers.get("Authorization", "")
if auth.startswith("Bearer "):
return auth[7:].strip()
return None
def _get_model_name(api_key: str) -> str:
"""Look up agent name for display as model name."""
try:
mongo = MongoDB.get_client()
db = mongo[settings.MONGO_DB_NAME]
agent = db["agents"].find_one({"key": api_key})
if agent:
return agent.get("name", api_key)
except Exception:
pass
return api_key
class _V1AnswerHelper(BaseAnswerResource):
"""Thin wrapper to access complete_stream / process_response_stream."""
pass
@v1_bp.route("/chat/completions", methods=["POST"])
def chat_completions():
"""Handle POST /v1/chat/completions."""
api_key = _extract_bearer_token()
if not api_key:
return make_response(
jsonify({"error": {"message": "Missing Authorization header", "type": "auth_error"}}),
401,
)
data = request.get_json()
if not data or not data.get("messages"):
return make_response(
jsonify({"error": {"message": "messages field is required", "type": "invalid_request"}}),
400,
)
is_stream = data.get("stream", False)
model_name = _get_model_name(api_key)
try:
internal_data = translate_request(data, api_key)
except Exception as e:
logger.error(f"/v1/chat/completions translate error: {e}", exc_info=True)
return make_response(
jsonify({"error": {"message": "Failed to process request", "type": "invalid_request"}}),
400,
)
# Use the api_key as decoded token for agent auth
decoded_token = {"sub": "api_key_user"}
try:
processor = StreamProcessor(internal_data, decoded_token)
if internal_data.get("tool_actions"):
# Continuation mode
conversation_id = internal_data.get("conversation_id")
if not conversation_id:
return make_response(
jsonify({"error": {"message": "conversation_id required for tool continuation", "type": "invalid_request"}}),
400,
)
(
agent,
messages,
tools_dict,
pending_tool_calls,
tool_actions,
) = processor.resume_from_tool_actions(
internal_data["tool_actions"], conversation_id
)
continuation = {
"messages": messages,
"tools_dict": tools_dict,
"pending_tool_calls": pending_tool_calls,
"tool_actions": tool_actions,
}
question = ""
else:
# Normal mode
question = internal_data.get("question", "")
agent = processor.build_agent(question)
continuation = None
if not processor.decoded_token:
return make_response(
jsonify({"error": {"message": "Unauthorized", "type": "auth_error"}}),
401,
)
helper = _V1AnswerHelper()
usage_error = helper.check_usage(processor.agent_config)
if usage_error:
return usage_error
if is_stream:
return Response(
_stream_response(
helper, question, agent, processor, model_name, continuation
),
mimetype="text/event-stream",
headers={
"Cache-Control": "no-cache",
"X-Accel-Buffering": "no",
},
)
else:
return _non_stream_response(
helper, question, agent, processor, model_name, continuation
)
except ValueError as e:
logger.error(
f"/v1/chat/completions error: {e} - {traceback.format_exc()}",
extra={"error": str(e)},
)
return make_response(
jsonify({"error": {"message": str(e), "type": "invalid_request"}}),
400,
)
except Exception as e:
logger.error(
f"/v1/chat/completions error: {e} - {traceback.format_exc()}",
extra={"error": str(e)},
)
return make_response(
jsonify({"error": {"message": "Internal server error", "type": "server_error"}}),
500,
)
def _stream_response(
helper: _V1AnswerHelper,
question: str,
agent: Any,
processor: StreamProcessor,
model_name: str,
continuation: Optional[Dict],
) -> Generator[str, None, None]:
"""Generate translated SSE chunks for streaming response."""
completion_id = f"chatcmpl-{int(time.time())}"
internal_stream = helper.complete_stream(
question=question,
agent=agent,
conversation_id=processor.conversation_id,
user_api_key=processor.agent_config.get("user_api_key"),
decoded_token=processor.decoded_token,
agent_id=processor.agent_id,
model_id=processor.model_id,
_continuation=continuation,
)
for line in internal_stream:
if not line.strip():
continue
# Parse the internal SSE event
event_str = line.replace("data: ", "").strip()
try:
event_data = json.loads(event_str)
except (json.JSONDecodeError, TypeError):
continue
# Update completion_id when we get the conversation id
if event_data.get("type") == "id":
conv_id = event_data.get("id", "")
if conv_id:
completion_id = f"chatcmpl-{conv_id}"
# Translate to standard format
translated = translate_stream_event(event_data, completion_id, model_name)
for chunk in translated:
yield chunk
def _non_stream_response(
helper: _V1AnswerHelper,
question: str,
agent: Any,
processor: StreamProcessor,
model_name: str,
continuation: Optional[Dict],
) -> Response:
"""Collect full response and return as single JSON."""
stream = helper.complete_stream(
question=question,
agent=agent,
conversation_id=processor.conversation_id,
user_api_key=processor.agent_config.get("user_api_key"),
decoded_token=processor.decoded_token,
agent_id=processor.agent_id,
model_id=processor.model_id,
_continuation=continuation,
)
result = helper.process_response_stream(stream)
if len(result) == 7:
conversation_id, answer, sources, tool_calls, thought, error, extra = result
else:
conversation_id, answer, sources, tool_calls, thought, error = result
extra = None
if error:
return make_response(
jsonify({"error": {"message": error, "type": "server_error"}}),
500,
)
pending = extra.get("pending_tool_calls") if isinstance(extra, dict) else None
response = translate_response(
conversation_id=conversation_id,
answer=answer or "",
sources=sources,
tool_calls=tool_calls,
thought=thought or "",
model_name=model_name,
pending_tool_calls=pending,
)
return make_response(jsonify(response), 200)
@v1_bp.route("/models", methods=["GET"])
def list_models():
"""Handle GET /v1/models — return agents as models."""
api_key = _extract_bearer_token()
if not api_key:
return make_response(
jsonify({"error": {"message": "Missing Authorization header", "type": "auth_error"}}),
401,
)
try:
mongo = MongoDB.get_client()
db = mongo[settings.MONGO_DB_NAME]
agents_collection = db["agents"]
# Find the agent for this api_key
agent = agents_collection.find_one({"key": api_key})
if not agent:
return make_response(
jsonify({"error": {"message": "Invalid API key", "type": "auth_error"}}),
401,
)
user = agent.get("user")
# Return all agents belonging to this user
user_agents = list(agents_collection.find({"user": user}))
models = []
for ag in user_agents:
created = ag.get("createdAt")
created_ts = int(created.timestamp()) if created else int(time.time())
models.append({
"id": str(ag.get("key", "")),
"object": "model",
"created": created_ts,
"owned_by": "docsgpt",
"name": ag.get("name", ""),
"description": ag.get("description", ""),
})
return make_response(
jsonify({"object": "list", "data": models}),
200,
)
except Exception as e:
logger.error(f"/v1/models error: {e}", exc_info=True)
return make_response(
jsonify({"error": {"message": "Internal server error", "type": "server_error"}}),
500,
)
+404
View File
@@ -0,0 +1,404 @@
"""Translate between standard chat completions format and DocsGPT internals.
This module handles:
- Request translation (chat completions -> DocsGPT internal format)
- Response translation (DocsGPT response -> chat completions format)
- Streaming event translation (DocsGPT SSE -> standard SSE chunks)
"""
import json
import time
from typing import Any, Dict, List, Optional
# ---------------------------------------------------------------------------
# Request translation
# ---------------------------------------------------------------------------
def is_continuation(messages: List[Dict]) -> bool:
"""Check if messages represent a tool-call continuation.
A continuation is detected when the last message(s) have ``role: "tool"``
immediately after an assistant message with ``tool_calls``.
"""
if not messages:
return False
# Walk backwards: if we see tool messages before hitting a non-tool, non-assistant message
# and there's an assistant message with tool_calls, it's a continuation.
i = len(messages) - 1
while i >= 0 and messages[i].get("role") == "tool":
i -= 1
if i < 0:
return False
return (
messages[i].get("role") == "assistant"
and bool(messages[i].get("tool_calls"))
)
def extract_tool_results(messages: List[Dict]) -> List[Dict]:
"""Extract tool results from trailing tool messages for continuation.
Returns a list of ``tool_actions`` dicts with ``call_id`` and ``result``.
"""
results = []
for msg in reversed(messages):
if msg.get("role") != "tool":
break
call_id = msg.get("tool_call_id", "")
content = msg.get("content", "")
if isinstance(content, str):
try:
content = json.loads(content)
except (json.JSONDecodeError, TypeError):
pass
results.append({"call_id": call_id, "result": content})
results.reverse()
return results
def extract_conversation_id(messages: List[Dict]) -> Optional[str]:
"""Try to extract conversation_id from the assistant message before tool results.
The conversation_id may be stored in a custom field on the assistant message
from a previous response cycle.
"""
for msg in reversed(messages):
if msg.get("role") == "assistant":
# Check docsgpt extension
return msg.get("docsgpt", {}).get("conversation_id")
return None
def convert_history(messages: List[Dict]) -> List[Dict]:
"""Convert chat completions messages array to DocsGPT history format.
DocsGPT history is a list of ``{prompt, response}`` dicts.
Excludes the last user message (that becomes the ``question``).
"""
history = []
i = 0
while i < len(messages):
msg = messages[i]
if msg.get("role") == "system":
i += 1
continue
if msg.get("role") == "user":
# Look ahead for assistant response
if i + 1 < len(messages) and messages[i + 1].get("role") == "assistant":
content = messages[i + 1].get("content") or ""
history.append({
"prompt": msg.get("content", ""),
"response": content,
})
i += 2
continue
# Last user message without response — skip (it's the question)
i += 1
continue
i += 1
return history
def translate_request(
data: Dict[str, Any], api_key: str
) -> Dict[str, Any]:
"""Translate a chat completions request to DocsGPT internal format.
Args:
data: The incoming request body.
api_key: Agent API key from the Authorization header.
Returns:
Dict suitable for passing to ``StreamProcessor``.
"""
messages = data.get("messages", [])
# Check for continuation (tool results after assistant tool_calls)
if is_continuation(messages):
tool_actions = extract_tool_results(messages)
conversation_id = extract_conversation_id(messages)
result = {
"conversation_id": conversation_id,
"tool_actions": tool_actions,
"api_key": api_key,
}
# Carry tools forward for next iteration
if data.get("tools"):
result["client_tools"] = data["tools"]
return result
# Normal request — extract question from last user message
question = ""
for msg in reversed(messages):
if msg.get("role") == "user":
question = msg.get("content", "")
break
history = convert_history(messages)
result = {
"question": question,
"api_key": api_key,
"history": json.dumps(history),
"save_conversation": True,
}
# Client tools (Phase 2)
if data.get("tools"):
result["client_tools"] = data["tools"]
# DocsGPT extensions
docsgpt = data.get("docsgpt", {})
if docsgpt.get("attachments"):
result["attachments"] = docsgpt["attachments"]
return result
# ---------------------------------------------------------------------------
# Response translation (non-streaming)
# ---------------------------------------------------------------------------
def translate_response(
conversation_id: str,
answer: str,
sources: Optional[List[Dict]],
tool_calls: Optional[List[Dict]],
thought: str,
model_name: str,
pending_tool_calls: Optional[List[Dict]] = None,
) -> Dict[str, Any]:
"""Translate DocsGPT response to chat completions format.
Args:
conversation_id: The DocsGPT conversation ID.
answer: The assistant's text response.
sources: RAG retrieval sources.
tool_calls: Completed tool call results.
thought: Reasoning/thinking tokens.
model_name: Model/agent identifier.
pending_tool_calls: Pending client-side tool calls (if paused).
Returns:
Dict in the standard chat completions response format.
"""
created = int(time.time())
completion_id = f"chatcmpl-{conversation_id}" if conversation_id else f"chatcmpl-{created}"
# Build message
message: Dict[str, Any] = {"role": "assistant"}
if pending_tool_calls:
# Tool calls pending — return them for client execution
message["content"] = None
message["tool_calls"] = [
{
"id": tc.get("call_id", ""),
"type": "function",
"function": {
"name": tc.get("name", tc.get("action_name", "")),
"arguments": (
json.dumps(tc["arguments"])
if isinstance(tc.get("arguments"), dict)
else tc.get("arguments", "{}")
),
},
}
for tc in pending_tool_calls
]
finish_reason = "tool_calls"
else:
message["content"] = answer
if thought:
message["reasoning_content"] = thought
finish_reason = "stop"
result: Dict[str, Any] = {
"id": completion_id,
"object": "chat.completion",
"created": created,
"model": model_name,
"choices": [
{
"index": 0,
"message": message,
"finish_reason": finish_reason,
}
],
"usage": {
"prompt_tokens": 0,
"completion_tokens": 0,
"total_tokens": 0,
},
}
# DocsGPT extensions
docsgpt: Dict[str, Any] = {}
if conversation_id:
docsgpt["conversation_id"] = conversation_id
if sources:
docsgpt["sources"] = sources
if tool_calls:
docsgpt["tool_calls"] = tool_calls
if docsgpt:
result["docsgpt"] = docsgpt
return result
# ---------------------------------------------------------------------------
# Streaming event translation
# ---------------------------------------------------------------------------
def _make_chunk(
completion_id: str,
model_name: str,
delta: Dict[str, Any],
finish_reason: Optional[str] = None,
) -> str:
"""Build a single SSE chunk in the standard streaming format."""
chunk = {
"id": completion_id,
"object": "chat.completion.chunk",
"created": int(time.time()),
"model": model_name,
"choices": [
{
"index": 0,
"delta": delta,
"finish_reason": finish_reason,
}
],
}
return f"data: {json.dumps(chunk)}\n\n"
def _make_docsgpt_chunk(data: Dict[str, Any]) -> str:
"""Build a DocsGPT extension SSE chunk."""
return f"data: {json.dumps({'docsgpt': data})}\n\n"
def translate_stream_event(
event_data: Dict[str, Any],
completion_id: str,
model_name: str,
) -> List[str]:
"""Translate a DocsGPT SSE event dict to standard streaming chunks.
May return 0, 1, or 2 chunks per input event. For example, a completed
tool call produces both a docsgpt extension chunk and nothing on the
standard side (since server-side tool calls aren't surfaced in standard
format).
Args:
event_data: Parsed DocsGPT event dict.
completion_id: The completion ID for this response.
model_name: Model/agent identifier.
Returns:
List of SSE-formatted strings to send to the client.
"""
event_type = event_data.get("type")
chunks: List[str] = []
if event_type == "answer":
chunks.append(
_make_chunk(completion_id, model_name, {"content": event_data.get("answer", "")})
)
elif event_type == "thought":
chunks.append(
_make_chunk(
completion_id, model_name,
{"reasoning_content": event_data.get("thought", "")},
)
)
elif event_type == "source":
chunks.append(
_make_docsgpt_chunk({
"type": "source",
"sources": event_data.get("source", []),
})
)
elif event_type == "tool_call":
tc_data = event_data.get("data", {})
status = tc_data.get("status")
if status == "requires_client_execution":
# Standard: stream as tool_calls delta
args = tc_data.get("arguments", {})
args_str = json.dumps(args) if isinstance(args, dict) else str(args)
chunks.append(
_make_chunk(completion_id, model_name, {
"tool_calls": [{
"index": 0,
"id": tc_data.get("call_id", ""),
"type": "function",
"function": {
"name": tc_data.get("action_name", ""),
"arguments": args_str,
},
}],
})
)
elif status == "awaiting_approval":
# Extension: approval needed
chunks.append(_make_docsgpt_chunk({"type": "tool_call", "data": tc_data}))
elif status in ("completed", "pending", "error", "denied", "skipped"):
# Extension: tool call progress
chunks.append(_make_docsgpt_chunk({"type": "tool_call", "data": tc_data}))
elif event_type == "tool_calls_pending":
# Standard: finish_reason = tool_calls
chunks.append(
_make_chunk(completion_id, model_name, {}, finish_reason="tool_calls")
)
# Also emit as docsgpt extension
chunks.append(
_make_docsgpt_chunk({
"type": "tool_calls_pending",
"pending_tool_calls": event_data.get("data", {}).get("pending_tool_calls", []),
})
)
elif event_type == "end":
chunks.append(
_make_chunk(completion_id, model_name, {}, finish_reason="stop")
)
chunks.append("data: [DONE]\n\n")
elif event_type == "id":
chunks.append(
_make_docsgpt_chunk({
"type": "id",
"conversation_id": event_data.get("id", ""),
})
)
elif event_type == "error":
# Emit as standard error (non-standard but widely supported)
error_data = {
"error": {
"message": event_data.get("error", "An error occurred"),
"type": "server_error",
}
}
chunks.append(f"data: {json.dumps(error_data)}\n\n")
elif event_type == "structured_answer":
chunks.append(
_make_chunk(
completion_id, model_name,
{"content": event_data.get("answer", "")},
)
)
# Skip: tool_calls (redundant), research_plan, research_progress
return chunks
+2
View File
@@ -17,6 +17,7 @@ from application.api.answer import answer # noqa: E402
from application.api.internal.routes import internal # noqa: E402
from application.api.user.routes import user # noqa: E402
from application.api.connector.routes import connector # noqa: E402
from application.api.v1 import v1_bp # noqa: E402
from application.celery_init import celery # noqa: E402
from application.core.settings import settings # noqa: E402
from application.stt.upload_limits import ( # noqa: E402
@@ -36,6 +37,7 @@ app.register_blueprint(user)
app.register_blueprint(answer)
app.register_blueprint(internal)
app.register_blueprint(connector)
app.register_blueprint(v1_bp)
app.config.update(
UPLOAD_FOLDER="inputs",
CELERY_BROKER_URL=settings.CELERY_BROKER_URL,
@@ -22,6 +22,7 @@ import {
resendQuery,
selectQueries,
selectStatus,
submitToolActions,
updateQuery,
} from './conversationSlice';
import { selectCompletedAttachments } from '../upload/uploadSlice';
@@ -41,6 +42,17 @@ export default function Conversation() {
const [lastQueryReturnedErr, setLastQueryReturnedErr] =
useState<boolean>(false);
const handleToolAction = useCallback(
(callId: string, decision: 'approved' | 'denied', comment?: string) => {
dispatch(
submitToolActions({
toolActions: [{ call_id: callId, decision, comment }],
}),
);
},
[dispatch],
);
const lastAutoOpenedArtifactId = useRef<string | null>(null);
const didInitArtifactAutoOpen = useRef(false);
const prevConversationId = useRef<string | null>(conversationId);
@@ -233,6 +245,7 @@ export default function Conversation() {
status={status}
showHeroOnEmpty={selectedAgent ? false : true}
onOpenArtifact={handleOpenArtifact}
onToolAction={handleToolAction}
isSplitView={isSplitArtifactOpen}
headerContent={
selectedAgent ? (
@@ -65,6 +65,11 @@ const ConversationBubble = forwardRef<
) => void;
filesAttached?: { id: string; fileName: string }[];
onOpenArtifact?: (artifact: { id: string; toolName: string }) => void;
onToolAction?: (
callId: string,
decision: 'approved' | 'denied',
comment?: string,
) => void;
}
>(function ConversationBubble(
{
@@ -83,6 +88,7 @@ const ConversationBubble = forwardRef<
handleUpdatedQuestionSubmission,
filesAttached,
onOpenArtifact,
onToolAction,
},
ref,
) {
@@ -411,7 +417,7 @@ const ConversationBubble = forwardRef<
)}
{research && <ResearchProgress research={research} />}
{toolCalls && toolCalls.length > 0 && (
<ToolCalls toolCalls={toolCalls} />
<ToolCalls toolCalls={toolCalls} onToolAction={onToolAction} />
)}
{!message && primaryArtifactCall?.artifact_id && onOpenArtifact && (
<div className="my-2 ml-2 flex justify-start">
@@ -884,8 +890,23 @@ function AllSources(sources: AllSourcesProps) {
}
export default ConversationBubble;
function ToolCalls({ toolCalls }: { toolCalls: ToolCallsType[] }) {
function ToolCalls({
toolCalls,
onToolAction,
}: {
toolCalls: ToolCallsType[];
onToolAction?: (
callId: string,
decision: 'approved' | 'denied',
comment?: string,
) => void;
}) {
const [isToolCallsOpen, setIsToolCallsOpen] = useState(false);
const [denyComments, setDenyComments] = useState<Record<string, string>>({});
const hasAwaitingApproval = toolCalls.some(
(tc) => tc.status === 'awaiting_approval',
);
return (
<div className="mb-4 flex w-full flex-col flex-wrap items-start self-start lg:flex-nowrap">
@@ -904,15 +925,22 @@ function ToolCalls({ toolCalls }: { toolCalls: ToolCallsType[] }) {
className="flex flex-row items-center gap-2"
onClick={() => setIsToolCallsOpen(!isToolCallsOpen)}
>
<p className="text-base font-semibold">Tool Calls</p>
<p className="text-base font-semibold">
Tool Calls
{hasAwaitingApproval && (
<span className="ml-2 text-xs font-normal text-yellow-600 dark:text-yellow-400">
(approval needed)
</span>
)}
</p>
<img
src={ChevronDown}
alt="ChevronDown"
className={`h-4 w-4 transform transition-transform duration-200 dark:invert ${isToolCallsOpen ? 'rotate-180' : ''}`}
className={`h-4 w-4 transform transition-transform duration-200 dark:invert ${isToolCallsOpen || hasAwaitingApproval ? 'rotate-180' : ''}`}
/>
</button>
</div>
{isToolCallsOpen && (
{(isToolCallsOpen || hasAwaitingApproval) && (
<div className="fade-in mr-5 ml-3 w-[90vw] md:w-[70vw] lg:w-full">
<div className="grid grid-cols-1 gap-2">
{toolCalls.map((toolCall, index) => (
@@ -921,6 +949,7 @@ function ToolCalls({ toolCalls }: { toolCalls: ToolCallsType[] }) {
title={`${toolCall.tool_name} - ${toolCall.action_name.substring(0, toolCall.action_name.lastIndexOf('_'))}`}
className="bg-muted dark:bg-answer-bubble w-full rounded-4xl"
titleClassName="px-6 py-2 text-sm font-semibold"
open={toolCall.status === 'awaiting_approval'}
>
<div className="flex flex-col gap-1">
<div className="border-border flex flex-col rounded-2xl border">
@@ -979,6 +1008,57 @@ function ToolCalls({ toolCalls }: { toolCalls: ToolCallsType[] }) {
</span>
</p>
)}
{toolCall.status === 'awaiting_approval' && (
<div className="dark:bg-card flex flex-col gap-2 rounded-b-2xl p-3">
<p className="text-sm text-yellow-600 dark:text-yellow-400">
This tool requires your approval before executing.
</p>
<input
type="text"
placeholder="Optional comment (for deny)..."
className="border-border bg-background w-full rounded-lg border px-3 py-1.5 text-sm"
value={denyComments[toolCall.call_id] || ''}
onChange={(e) =>
setDenyComments((prev) => ({
...prev,
[toolCall.call_id]: e.target.value,
}))
}
/>
<div className="flex gap-2">
<button
className="rounded-lg bg-green-600 px-4 py-1.5 text-sm font-medium text-white hover:bg-green-700"
onClick={() =>
onToolAction?.(toolCall.call_id, 'approved')
}
>
Approve
</button>
<button
className="rounded-lg bg-red-600 px-4 py-1.5 text-sm font-medium text-white hover:bg-red-700"
onClick={() =>
onToolAction?.(
toolCall.call_id,
'denied',
denyComments[toolCall.call_id],
)
}
>
Deny
</button>
</div>
</div>
)}
{toolCall.status === 'denied' && (
<p className="dark:bg-card rounded-b-2xl p-2 font-mono text-sm wrap-break-word">
<span
className="leading-[23px] text-orange-500 dark:text-orange-400"
style={{ fontFamily: 'IBMPlexMono-Medium' }}
>
Denied by user
</span>
</p>
)}
</div>
</div>
</Accordion>
@@ -38,6 +38,11 @@ type ConversationMessagesProps = {
showHeroOnEmpty?: boolean;
headerContent?: ReactNode;
onOpenArtifact?: (artifact: { id: string; toolName: string }) => void;
onToolAction?: (
callId: string,
decision: 'approved' | 'denied',
comment?: string,
) => void;
isSplitView?: boolean;
};
@@ -50,6 +55,7 @@ export default function ConversationMessages({
showHeroOnEmpty = true,
headerContent,
onOpenArtifact,
onToolAction,
isSplitView = false,
}: ConversationMessagesProps) {
const [isDarkTheme] = useDarkTheme();
@@ -154,6 +160,7 @@ export default function ConversationMessages({
toolCalls={query.tool_calls}
research={query.research}
onOpenArtifact={onOpenArtifact}
onToolAction={onToolAction}
feedback={query.feedback}
isStreaming={isCurrentlyStreaming}
handleFeedback={
@@ -188,6 +188,72 @@ export function handleFetchAnswerSteaming(
});
}
export function handleSubmitToolActions(
conversationId: string,
toolActions: {
call_id: string;
decision?: 'approved' | 'denied';
comment?: string;
result?: Record<string, any>;
}[],
token: string | null,
signal: AbortSignal,
onEvent: (event: MessageEvent) => void,
): Promise<Answer> {
const payload = {
conversation_id: conversationId,
tool_actions: toolActions,
};
return new Promise<Answer>((resolve, reject) => {
conversationService
.answerStream(payload, token, signal)
.then((response) => {
if (!response.body) throw Error('No response body');
let buffer = '';
const reader = response.body.getReader();
const decoder = new TextDecoder('utf-8');
const processStream = ({
done,
value,
}: ReadableStreamReadResult<Uint8Array>) => {
if (done) return;
const chunk = decoder.decode(value);
buffer += chunk;
const events = buffer.split('\n\n');
buffer = events.pop() ?? '';
for (const event of events) {
if (event.trim().startsWith('data:')) {
const dataLine: string = event
.split('\n')
.map((line: string) => line.replace(/^data:\s?/, ''))
.join('');
const messageEvent = new MessageEvent('message', {
data: dataLine.trim(),
});
onEvent(messageEvent);
}
}
reader.read().then(processStream).catch(reject);
};
reader.read().then(processStream).catch(reject);
})
.catch((error) => {
console.error('Tool actions submission failed:', error);
reject(error);
});
});
}
export function handleSearch(
question: string,
token: string | null,
@@ -1,7 +1,7 @@
import { ToolCallsType } from './types';
export type MESSAGE_TYPE = 'QUESTION' | 'ANSWER' | 'ERROR';
export type Status = 'idle' | 'loading' | 'failed';
export type Status = 'idle' | 'loading' | 'failed' | 'awaiting_tool_actions';
export type FEEDBACK = 'LIKE' | 'DISLIKE' | null;
export interface Message {
@@ -10,6 +10,7 @@ import {
import {
handleFetchAnswer,
handleFetchAnswerSteaming,
handleSubmitToolActions,
} from './conversationHandlers';
import {
Answer,
@@ -138,6 +139,10 @@ export const fetchAnswer = createAsyncThunk<
tool_call: data.data as ToolCallsType,
}),
);
} else if (data.type === 'tool_calls_pending') {
dispatch(
conversationSlice.actions.setStatus('awaiting_tool_actions'),
);
} else if (data.type === 'research_plan') {
dispatch(
updateResearchPlan({
@@ -260,6 +265,94 @@ export const fetchAnswer = createAsyncThunk<
};
});
export const submitToolActions = createAsyncThunk<
void,
{
toolActions: {
call_id: string;
decision?: 'approved' | 'denied';
comment?: string;
result?: Record<string, any>;
}[];
}
>('submitToolActions', async ({ toolActions }, { dispatch, getState }) => {
if (abortController) abortController.abort();
abortController = new AbortController();
const { signal } = abortController;
const state = getState() as RootState;
const conversationId = state.conversation.conversationId;
if (!conversationId) return;
dispatch(conversationSlice.actions.setStatus('loading'));
await handleSubmitToolActions(
conversationId,
toolActions,
state.preference.token,
signal,
(event) => {
const data = JSON.parse(event.data);
const targetIndex = state.conversation.queries.length - 1;
if (data.type === 'end') {
dispatch(conversationSlice.actions.setStatus('idle'));
getConversations(state.preference.token)
.then((fetchedConversations) => {
dispatch(setConversations(fetchedConversations));
})
.catch((error) => {
console.error('Failed to fetch conversations: ', error);
});
} else if (data.type === 'id') {
// conversation ID already set
} else if (data.type === 'thought') {
dispatch(
updateThought({
conversationId,
index: targetIndex,
query: { thought: data.thought },
}),
);
} else if (data.type === 'source') {
dispatch(
updateStreamingSource({
conversationId,
index: targetIndex,
query: { sources: data.source ?? [] },
}),
);
} else if (data.type === 'tool_call') {
dispatch(
updateToolCall({
index: targetIndex,
tool_call: data.data as ToolCallsType,
}),
);
} else if (data.type === 'tool_calls_pending') {
dispatch(conversationSlice.actions.setStatus('awaiting_tool_actions'));
} else if (data.type === 'error') {
dispatch(conversationSlice.actions.setStatus('failed'));
dispatch(
conversationSlice.actions.raiseError({
conversationId,
index: targetIndex,
message: data.error,
}),
);
} else if (data.type === 'answer') {
dispatch(
updateStreamingQuery({
conversationId,
index: targetIndex,
query: { response: data.answer },
}),
);
}
},
);
});
export const conversationSlice = createSlice({
name: 'conversation',
initialState,
+7 -1
View File
@@ -5,6 +5,12 @@ export type ToolCallsType = {
arguments: Record<string, any>;
result?: Record<string, any>;
error?: string;
status?: 'pending' | 'completed' | 'error';
status?:
| 'pending'
| 'completed'
| 'error'
| 'awaiting_approval'
| 'denied'
| 'requires_client_execution';
artifact_id?: string;
};
+54 -1
View File
@@ -487,9 +487,33 @@ export default function ToolConfig({
)}
</div>
<div
className="flex items-center gap-2"
className="flex items-center gap-3"
onClick={(e) => e.stopPropagation()}
>
<div className="flex items-center gap-1">
<span className="text-xs text-gray-500 dark:text-gray-400">
{t('settings.tools.requireApproval', 'Approval')}
</span>
<ToggleSwitch
checked={action.require_approval ?? false}
onChange={(checked) => {
setTool({
...tool,
actions: tool.actions.map((act, index) => {
if (index === originalIndex) {
return {
...act,
require_approval: checked,
};
}
return act;
}),
});
}}
size="small"
id={`approvalToggle-${originalIndex}`}
/>
</div>
<ToggleSwitch
checked={action.active}
onChange={(checked) => {
@@ -926,6 +950,35 @@ function APIToolConfig({
className="h-4 w-4 opacity-40 transition-opacity hover:opacity-100"
/>
</button>
<div className="flex items-center gap-1">
<span className="text-xs text-gray-500 dark:text-gray-400">
{t('settings.tools.requireApproval', 'Approval')}
</span>
<ToggleSwitch
checked={action.require_approval ?? false}
onChange={() => {
setApiTool((prevApiTool) => {
const updatedActions = {
...prevApiTool.config.actions,
};
updatedActions[actionName] = {
...updatedActions[actionName],
require_approval:
!updatedActions[actionName].require_approval,
};
return {
...prevApiTool,
config: {
...prevApiTool.config,
actions: updatedActions,
},
};
});
}}
size="small"
id={`approvalToggle-${actionIndex}`}
/>
</div>
<ToggleSwitch
checked={action.active}
onChange={() => handleActionToggle(actionName)}
+2
View File
@@ -69,6 +69,7 @@ export type UserToolType = {
type: string;
};
active: boolean;
require_approval?: boolean;
}[];
};
@@ -81,6 +82,7 @@ export type APIActionType = {
headers: ParameterGroupType;
body: ParameterGroupType;
active: boolean;
require_approval?: boolean;
body_content_type?:
| 'application/json'
| 'application/x-www-form-urlencoded'
+481
View File
@@ -0,0 +1,481 @@
"""Tests for tool approval (Phase 3).
Covers require_approval flag, check_pause for approval, the handler
pause/resume flow, and gen_continuation with approved/denied actions.
"""
from unittest.mock import Mock
import pytest
from application.agents.tool_executor import ToolExecutor
from application.llm.handlers.base import LLMHandler, LLMResponse, ToolCall
# ---------------------------------------------------------------------------
# check_pause with require_approval
# ---------------------------------------------------------------------------
@pytest.mark.unit
class TestCheckPauseApproval:
def _make_call(self, name="action_0", call_id="c1"):
call = Mock()
call.name = name
call.id = call_id
call.arguments = "{}"
call.thought_signature = None
return call
def test_approval_required_triggers_pause(self):
executor = ToolExecutor()
tools_dict = {
"0": {
"name": "telegram",
"actions": [
{
"name": "send_msg",
"active": True,
"require_approval": True,
"parameters": {},
},
],
}
}
call = self._make_call(name="send_msg_0")
result = executor.check_pause(tools_dict, call, "OpenAILLM")
assert result is not None
assert result["pause_type"] == "awaiting_approval"
assert result["tool_name"] == "telegram"
assert result["action_name"] == "send_msg"
assert result["tool_id"] == "0"
def test_approval_not_required_no_pause(self):
executor = ToolExecutor()
tools_dict = {
"0": {
"name": "brave",
"actions": [
{
"name": "search",
"active": True,
"require_approval": False,
"parameters": {},
},
],
}
}
call = self._make_call(name="search_0")
result = executor.check_pause(tools_dict, call, "OpenAILLM")
assert result is None
def test_approval_absent_defaults_to_false(self):
executor = ToolExecutor()
tools_dict = {
"0": {
"name": "brave",
"actions": [
{
"name": "search",
"active": True,
"parameters": {},
},
],
}
}
call = self._make_call(name="search_0")
result = executor.check_pause(tools_dict, call, "OpenAILLM")
assert result is None
def test_api_tool_approval(self):
executor = ToolExecutor()
tools_dict = {
"0": {
"name": "api_tool",
"config": {
"actions": {
"delete_user": {
"name": "delete_user",
"require_approval": True,
"url": "http://example.com",
"method": "DELETE",
"active": True,
}
}
},
}
}
call = self._make_call(name="delete_user_0")
result = executor.check_pause(tools_dict, call, "OpenAILLM")
assert result is not None
assert result["pause_type"] == "awaiting_approval"
def test_api_tool_no_approval(self):
executor = ToolExecutor()
tools_dict = {
"0": {
"name": "api_tool",
"config": {
"actions": {
"list_users": {
"name": "list_users",
"url": "http://example.com",
"method": "GET",
"active": True,
}
}
},
}
}
call = self._make_call(name="list_users_0")
result = executor.check_pause(tools_dict, call, "OpenAILLM")
assert result is None
# ---------------------------------------------------------------------------
# Handler: approval tool causes pause signal
# ---------------------------------------------------------------------------
class ConcreteHandler(LLMHandler):
def parse_response(self, response):
return LLMResponse(
content=str(response), tool_calls=[], finish_reason="stop",
raw_response=response,
)
def create_tool_message(self, tool_call, result):
import json as _json
content = _json.dumps(result) if not isinstance(result, str) else result
return {"role": "tool", "tool_call_id": tool_call.id, "content": content}
def _iterate_stream(self, response):
for chunk in response:
yield chunk
@pytest.mark.unit
class TestHandlerApprovalPause:
def _make_agent(self, pause_return):
agent = Mock()
agent._check_context_limit = Mock(return_value=False)
agent.context_limit_reached = False
agent.llm.__class__.__name__ = "MockLLM"
agent.tool_executor.check_pause = Mock(return_value=pause_return)
def fake_execute(tools_dict, call):
yield {"type": "tool_call", "data": {"status": "pending"}}
return ("tool result", call.id)
agent._execute_tool_action = Mock(side_effect=fake_execute)
return agent
def test_approval_tool_pauses(self):
handler = ConcreteHandler()
pause_info = {
"call_id": "c1",
"name": "send_msg_0",
"tool_name": "telegram",
"tool_id": "0",
"action_name": "send_msg",
"arguments": {"text": "hello"},
"pause_type": "awaiting_approval",
"thought_signature": None,
}
agent = self._make_agent(pause_info)
call = ToolCall(id="c1", name="send_msg_0", arguments='{"text": "hello"}')
gen = handler.handle_tool_calls(
agent, [call], {"0": {"name": "telegram"}}, []
)
events = []
pending = None
try:
while True:
events.append(next(gen))
except StopIteration as e:
messages, pending = e.value
assert pending is not None
assert len(pending) == 1
assert pending[0]["pause_type"] == "awaiting_approval"
# Should NOT have executed the tool
assert agent._execute_tool_action.call_count == 0
# Should have yielded awaiting_approval status
approval_events = [
e for e in events
if e.get("type") == "tool_call"
and e.get("data", {}).get("status") == "awaiting_approval"
]
assert len(approval_events) == 1
def test_mixed_normal_and_approval(self):
"""First tool runs normally, second needs approval."""
handler = ConcreteHandler()
call_count = {"n": 0}
def selective_pause(tools_dict, call, llm_class):
call_count["n"] += 1
if call_count["n"] == 2:
return {
"call_id": "c2",
"name": "send_msg_0",
"tool_name": "telegram",
"tool_id": "0",
"action_name": "send_msg",
"arguments": {},
"pause_type": "awaiting_approval",
"thought_signature": None,
}
return None
agent = self._make_agent(None)
agent.tool_executor.check_pause = Mock(side_effect=selective_pause)
calls = [
ToolCall(id="c1", name="search_0", arguments="{}"),
ToolCall(id="c2", name="send_msg_0", arguments="{}"),
]
gen = handler.handle_tool_calls(
agent, calls, {"0": {"name": "multi"}}, []
)
events = []
try:
while True:
events.append(next(gen))
except StopIteration as e:
messages, pending = e.value
# First tool executed
assert agent._execute_tool_action.call_count == 1
# Second tool is pending
assert pending is not None
assert len(pending) == 1
assert pending[0]["call_id"] == "c2"
# ---------------------------------------------------------------------------
# gen_continuation: approval and denial flows
# ---------------------------------------------------------------------------
@pytest.mark.unit
class TestGenContinuationApproval:
def _make_agent(self):
from application.agents.classic_agent import ClassicAgent
mock_llm = Mock()
mock_llm._supports_tools = True
mock_llm.gen_stream = Mock(return_value=iter(["Answer"]))
mock_llm._supports_structured_output = Mock(return_value=False)
mock_llm.__class__.__name__ = "MockLLM"
mock_handler = Mock()
mock_handler.process_message_flow = Mock(return_value=iter([]))
mock_handler.create_tool_message = Mock(
return_value={"role": "tool", "tool_call_id": "c1", "content": "result"}
)
mock_executor = Mock()
mock_executor.tool_calls = []
mock_executor.prepare_tools_for_llm = Mock(return_value=[])
mock_executor.get_truncated_tool_calls = Mock(return_value=[])
def fake_execute(tools_dict, call, llm_class):
yield {"type": "tool_call", "data": {"status": "pending"}}
return ("executed_result", "c1")
mock_executor.execute = Mock(side_effect=fake_execute)
agent = ClassicAgent(
endpoint="stream",
llm_name="openai",
model_id="gpt-4",
api_key="test",
llm=mock_llm,
llm_handler=mock_handler,
tool_executor=mock_executor,
)
return agent, mock_executor, mock_handler
def test_approved_tool_executes(self):
agent, mock_executor, mock_handler = self._make_agent()
messages = [{"role": "system", "content": "test"}]
pending = [
{
"call_id": "c1",
"name": "send_msg_0",
"tool_name": "telegram",
"tool_id": "0",
"action_name": "send_msg",
"arguments": {"text": "hello"},
"pause_type": "awaiting_approval",
"thought_signature": None,
}
]
tool_actions = [{"call_id": "c1", "decision": "approved"}]
list(agent.gen_continuation(
messages, {"0": {"name": "telegram"}}, pending, tool_actions
))
# Tool should have been executed
assert mock_executor.execute.called
def test_denied_tool_sends_denial_to_llm(self):
agent, mock_executor, mock_handler = self._make_agent()
messages = [{"role": "system", "content": "test"}]
pending = [
{
"call_id": "c1",
"name": "send_msg_0",
"tool_name": "telegram",
"tool_id": "0",
"action_name": "send_msg",
"arguments": {},
"pause_type": "awaiting_approval",
"thought_signature": None,
}
]
tool_actions = [
{"call_id": "c1", "decision": "denied", "comment": "not safe"},
]
events = list(agent.gen_continuation(
messages, {"0": {"name": "telegram"}}, pending, tool_actions
))
# Tool should NOT have been executed
assert not mock_executor.execute.called
# Should have a denied event
denied = [
e for e in events
if isinstance(e, dict)
and e.get("type") == "tool_call"
and e.get("data", {}).get("status") == "denied"
]
assert len(denied) == 1
# create_tool_message should have been called with denial text
denial_text = mock_handler.create_tool_message.call_args[0][1]
assert "denied" in denial_text.lower()
assert "not safe" in denial_text
def test_denied_without_comment(self):
agent, mock_executor, mock_handler = self._make_agent()
messages = [{"role": "system", "content": "test"}]
pending = [
{
"call_id": "c1",
"name": "act_0",
"tool_name": "tool",
"tool_id": "0",
"action_name": "act",
"arguments": {},
"pause_type": "awaiting_approval",
"thought_signature": None,
}
]
tool_actions = [{"call_id": "c1", "decision": "denied"}]
list(agent.gen_continuation(
messages, {"0": {"name": "tool"}}, pending, tool_actions
))
denial_text = mock_handler.create_tool_message.call_args[0][1]
assert "denied" in denial_text.lower()
def test_mixed_approve_deny_batch(self):
"""Two tools: one approved, one denied."""
agent, mock_executor, mock_handler = self._make_agent()
messages = [{"role": "system", "content": "test"}]
pending = [
{
"call_id": "c1",
"name": "safe_0",
"tool_name": "safe",
"tool_id": "0",
"action_name": "safe",
"arguments": {},
"pause_type": "awaiting_approval",
"thought_signature": None,
},
{
"call_id": "c2",
"name": "danger_0",
"tool_name": "danger",
"tool_id": "0",
"action_name": "danger",
"arguments": {},
"pause_type": "awaiting_approval",
"thought_signature": None,
},
]
tool_actions = [
{"call_id": "c1", "decision": "approved"},
{"call_id": "c2", "decision": "denied", "comment": "too risky"},
]
events = list(agent.gen_continuation(
messages, {"0": {"name": "multi"}}, pending, tool_actions
))
# First tool executed, second denied
assert mock_executor.execute.call_count == 1
denied = [
e for e in events
if isinstance(e, dict)
and e.get("type") == "tool_call"
and e.get("data", {}).get("status") == "denied"
]
assert len(denied) == 1
def test_missing_action_defaults_to_denial(self):
"""If client doesn't respond for a pending tool, treat as denied."""
agent, mock_executor, mock_handler = self._make_agent()
messages = [{"role": "system", "content": "test"}]
pending = [
{
"call_id": "c1",
"name": "act_0",
"tool_name": "tool",
"tool_id": "0",
"action_name": "act",
"arguments": {},
"pause_type": "awaiting_approval",
"thought_signature": None,
}
]
# Empty tool_actions — no response for c1
tool_actions = []
events = list(agent.gen_continuation(
messages, {"0": {"name": "tool"}}, pending, tool_actions
))
# Should have been treated as denied
assert not mock_executor.execute.called
denied = [
e for e in events
if isinstance(e, dict)
and e.get("type") == "tool_call"
and e.get("data", {}).get("status") == "denied"
]
assert len(denied) == 1
+459
View File
@@ -0,0 +1,459 @@
"""Tests for the v1 API translator (Phase 4).
Covers request translation, response translation, streaming event
translation, continuation detection, and history conversion.
"""
import json
import pytest
from application.api.v1.translator import (
convert_history,
extract_tool_results,
is_continuation,
translate_request,
translate_response,
translate_stream_event,
)
# ---------------------------------------------------------------------------
# is_continuation
# ---------------------------------------------------------------------------
@pytest.mark.unit
class TestIsContinuation:
def test_normal_messages_not_continuation(self):
messages = [
{"role": "user", "content": "Hello"},
]
assert is_continuation(messages) is False
def test_tool_after_assistant_tool_calls_is_continuation(self):
messages = [
{"role": "user", "content": "What's the weather?"},
{
"role": "assistant",
"content": None,
"tool_calls": [{"id": "c1", "type": "function", "function": {"name": "get_weather", "arguments": "{}"}}],
},
{"role": "tool", "tool_call_id": "c1", "content": '{"temp": "72F"}'},
]
assert is_continuation(messages) is True
def test_assistant_without_tool_calls_not_continuation(self):
messages = [
{"role": "user", "content": "Hello"},
{"role": "assistant", "content": "Hi"},
{"role": "tool", "tool_call_id": "c1", "content": "result"},
]
# assistant has no tool_calls — not a valid continuation
assert is_continuation(messages) is False
def test_empty_messages(self):
assert is_continuation([]) is False
def test_multiple_tool_results(self):
messages = [
{"role": "user", "content": "Do stuff"},
{
"role": "assistant",
"tool_calls": [
{"id": "c1", "type": "function", "function": {"name": "a", "arguments": "{}"}},
{"id": "c2", "type": "function", "function": {"name": "b", "arguments": "{}"}},
],
},
{"role": "tool", "tool_call_id": "c1", "content": "r1"},
{"role": "tool", "tool_call_id": "c2", "content": "r2"},
]
assert is_continuation(messages) is True
# ---------------------------------------------------------------------------
# extract_tool_results
# ---------------------------------------------------------------------------
@pytest.mark.unit
class TestExtractToolResults:
def test_extracts_results(self):
messages = [
{"role": "assistant", "tool_calls": [{"id": "c1"}]},
{"role": "tool", "tool_call_id": "c1", "content": '{"temp": "72F"}'},
]
results = extract_tool_results(messages)
assert len(results) == 1
assert results[0]["call_id"] == "c1"
assert results[0]["result"] == {"temp": "72F"}
def test_string_content(self):
messages = [
{"role": "tool", "tool_call_id": "c1", "content": "plain text"},
]
results = extract_tool_results(messages)
assert results[0]["result"] == "plain text"
def test_multiple_results(self):
messages = [
{"role": "assistant", "tool_calls": []},
{"role": "tool", "tool_call_id": "c1", "content": "r1"},
{"role": "tool", "tool_call_id": "c2", "content": "r2"},
]
results = extract_tool_results(messages)
assert len(results) == 2
assert results[0]["call_id"] == "c1"
assert results[1]["call_id"] == "c2"
# ---------------------------------------------------------------------------
# convert_history
# ---------------------------------------------------------------------------
@pytest.mark.unit
class TestConvertHistory:
def test_user_assistant_pairs(self):
messages = [
{"role": "system", "content": "You are helpful"},
{"role": "user", "content": "Hello"},
{"role": "assistant", "content": "Hi there"},
{"role": "user", "content": "How are you?"},
{"role": "assistant", "content": "I'm good"},
{"role": "user", "content": "What's 2+2?"}, # Last user = question
]
history = convert_history(messages)
assert len(history) == 2
assert history[0]["prompt"] == "Hello"
assert history[0]["response"] == "Hi there"
assert history[1]["prompt"] == "How are you?"
assert history[1]["response"] == "I'm good"
def test_single_user_message(self):
messages = [{"role": "user", "content": "Hi"}]
history = convert_history(messages)
assert history == []
def test_system_messages_skipped(self):
messages = [
{"role": "system", "content": "System prompt"},
{"role": "user", "content": "Question"},
]
history = convert_history(messages)
assert history == []
# ---------------------------------------------------------------------------
# translate_request
# ---------------------------------------------------------------------------
@pytest.mark.unit
class TestTranslateRequest:
def test_normal_request(self):
data = {
"messages": [
{"role": "user", "content": "Hello"},
{"role": "assistant", "content": "Hi"},
{"role": "user", "content": "What's 2+2?"},
],
}
result = translate_request(data, "test-key")
assert result["question"] == "What's 2+2?"
assert result["api_key"] == "test-key"
assert result["save_conversation"] is True
history = json.loads(result["history"])
assert len(history) == 1
assert history[0]["prompt"] == "Hello"
def test_continuation_request(self):
data = {
"messages": [
{"role": "user", "content": "Search for X"},
{
"role": "assistant",
"tool_calls": [{"id": "c1", "type": "function", "function": {"name": "search", "arguments": "{}"}}],
},
{"role": "tool", "tool_call_id": "c1", "content": '{"found": true}'},
],
}
result = translate_request(data, "key")
assert "tool_actions" in result
assert len(result["tool_actions"]) == 1
assert result["tool_actions"][0]["call_id"] == "c1"
def test_client_tools_passed_through(self):
data = {
"messages": [{"role": "user", "content": "Hi"}],
"tools": [{"type": "function", "function": {"name": "my_tool"}}],
}
result = translate_request(data, "key")
assert result["client_tools"] == data["tools"]
def test_docsgpt_attachments(self):
data = {
"messages": [{"role": "user", "content": "Hi"}],
"docsgpt": {"attachments": ["att1", "att2"]},
}
result = translate_request(data, "key")
assert result["attachments"] == ["att1", "att2"]
# ---------------------------------------------------------------------------
# translate_response
# ---------------------------------------------------------------------------
@pytest.mark.unit
class TestTranslateResponse:
def test_basic_response(self):
resp = translate_response(
conversation_id="conv-1",
answer="Hello!",
sources=[],
tool_calls=[],
thought="",
model_name="my-agent",
)
assert resp["id"] == "chatcmpl-conv-1"
assert resp["object"] == "chat.completion"
assert resp["model"] == "my-agent"
assert resp["choices"][0]["message"]["content"] == "Hello!"
assert resp["choices"][0]["finish_reason"] == "stop"
assert "reasoning_content" not in resp["choices"][0]["message"]
def test_response_with_thought(self):
resp = translate_response(
conversation_id="c1",
answer="Result",
sources=[],
tool_calls=[],
thought="Thinking about it...",
model_name="agent",
)
assert resp["choices"][0]["message"]["reasoning_content"] == "Thinking about it..."
def test_response_with_sources(self):
sources = [{"title": "doc.txt", "text": "content", "source": "/doc.txt"}]
resp = translate_response(
conversation_id="c1",
answer="Found it",
sources=sources,
tool_calls=[],
thought="",
model_name="agent",
)
assert resp["docsgpt"]["sources"] == sources
def test_response_with_tool_calls(self):
tool_calls = [{"tool_name": "notes", "call_id": "c1", "artifact_id": "a1"}]
resp = translate_response(
conversation_id="c1",
answer="Done",
sources=[],
tool_calls=tool_calls,
thought="",
model_name="agent",
)
assert resp["docsgpt"]["tool_calls"] == tool_calls
def test_pending_tool_calls(self):
pending = [
{
"call_id": "c1",
"name": "get_weather",
"arguments": {"city": "SF"},
}
]
resp = translate_response(
conversation_id="c1",
answer="",
sources=[],
tool_calls=[],
thought="",
model_name="agent",
pending_tool_calls=pending,
)
assert resp["choices"][0]["finish_reason"] == "tool_calls"
assert resp["choices"][0]["message"]["content"] is None
assert len(resp["choices"][0]["message"]["tool_calls"]) == 1
tc = resp["choices"][0]["message"]["tool_calls"][0]
assert tc["id"] == "c1"
assert tc["function"]["name"] == "get_weather"
def test_no_docsgpt_when_empty(self):
resp = translate_response(
conversation_id="",
answer="Hi",
sources=None,
tool_calls=None,
thought="",
model_name="agent",
)
assert "docsgpt" not in resp
# ---------------------------------------------------------------------------
# translate_stream_event
# ---------------------------------------------------------------------------
@pytest.mark.unit
class TestTranslateStreamEvent:
def test_answer_event(self):
chunks = translate_stream_event(
{"type": "answer", "answer": "Hello"},
"chatcmpl-1", "agent",
)
assert len(chunks) == 1
parsed = json.loads(chunks[0].replace("data: ", "").strip())
assert parsed["choices"][0]["delta"]["content"] == "Hello"
def test_thought_event(self):
chunks = translate_stream_event(
{"type": "thought", "thought": "reasoning"},
"chatcmpl-1", "agent",
)
assert len(chunks) == 1
parsed = json.loads(chunks[0].replace("data: ", "").strip())
assert parsed["choices"][0]["delta"]["reasoning_content"] == "reasoning"
def test_source_event(self):
chunks = translate_stream_event(
{"type": "source", "source": [{"title": "t", "text": "x"}]},
"chatcmpl-1", "agent",
)
assert len(chunks) == 1
parsed = json.loads(chunks[0].replace("data: ", "").strip())
assert parsed["docsgpt"]["type"] == "source"
assert len(parsed["docsgpt"]["sources"]) == 1
def test_end_event(self):
chunks = translate_stream_event(
{"type": "end"},
"chatcmpl-1", "agent",
)
assert len(chunks) == 2
# First chunk: finish_reason stop
parsed = json.loads(chunks[0].replace("data: ", "").strip())
assert parsed["choices"][0]["finish_reason"] == "stop"
# Second chunk: [DONE]
assert chunks[1].strip() == "data: [DONE]"
def test_tool_call_client_execution(self):
chunks = translate_stream_event(
{
"type": "tool_call",
"data": {
"call_id": "c1",
"action_name": "get_weather",
"arguments": {"city": "SF"},
"status": "requires_client_execution",
},
},
"chatcmpl-1", "agent",
)
assert len(chunks) == 1
parsed = json.loads(chunks[0].replace("data: ", "").strip())
tc = parsed["choices"][0]["delta"]["tool_calls"][0]
assert tc["id"] == "c1"
assert tc["function"]["name"] == "get_weather"
def test_tool_call_completed(self):
chunks = translate_stream_event(
{
"type": "tool_call",
"data": {
"call_id": "c1",
"status": "completed",
"result": "done",
"artifact_id": "a1",
},
},
"chatcmpl-1", "agent",
)
assert len(chunks) == 1
parsed = json.loads(chunks[0].replace("data: ", "").strip())
assert parsed["docsgpt"]["type"] == "tool_call"
assert parsed["docsgpt"]["data"]["artifact_id"] == "a1"
def test_tool_calls_pending(self):
chunks = translate_stream_event(
{
"type": "tool_calls_pending",
"data": {"pending_tool_calls": [{"call_id": "c1"}]},
},
"chatcmpl-1", "agent",
)
assert len(chunks) == 2
# Standard chunk with finish_reason tool_calls
parsed = json.loads(chunks[0].replace("data: ", "").strip())
assert parsed["choices"][0]["finish_reason"] == "tool_calls"
# Extension chunk
ext = json.loads(chunks[1].replace("data: ", "").strip())
assert ext["docsgpt"]["type"] == "tool_calls_pending"
def test_id_event(self):
chunks = translate_stream_event(
{"type": "id", "id": "conv-123"},
"chatcmpl-1", "agent",
)
assert len(chunks) == 1
parsed = json.loads(chunks[0].replace("data: ", "").strip())
assert parsed["docsgpt"]["conversation_id"] == "conv-123"
def test_error_event(self):
chunks = translate_stream_event(
{"type": "error", "error": "Something went wrong"},
"chatcmpl-1", "agent",
)
assert len(chunks) == 1
parsed = json.loads(chunks[0].replace("data: ", "").strip())
assert parsed["error"]["message"] == "Something went wrong"
def test_tool_calls_event_skipped(self):
"""The aggregate tool_calls event is redundant and should be skipped."""
chunks = translate_stream_event(
{"type": "tool_calls", "tool_calls": [{"call_id": "c1"}]},
"chatcmpl-1", "agent",
)
assert len(chunks) == 0
def test_research_events_skipped(self):
assert translate_stream_event(
{"type": "research_plan", "data": {}}, "id", "m"
) == []
assert translate_stream_event(
{"type": "research_progress", "data": {}}, "id", "m"
) == []
def test_awaiting_approval_as_extension(self):
chunks = translate_stream_event(
{
"type": "tool_call",
"data": {"call_id": "c1", "status": "awaiting_approval"},
},
"chatcmpl-1", "agent",
)
assert len(chunks) == 1
parsed = json.loads(chunks[0].replace("data: ", "").strip())
assert parsed["docsgpt"]["type"] == "tool_call"
def test_standard_clients_can_ignore_docsgpt(self):
"""Standard clients parse only 'choices' — docsgpt namespace is ignored."""
chunks = translate_stream_event(
{"type": "source", "source": [{"title": "t"}]},
"chatcmpl-1", "agent",
)
parsed = json.loads(chunks[0].replace("data: ", "").strip())
# No "choices" key — standard parsers skip this chunk entirely
assert "choices" not in parsed
# docsgpt key is present
assert "docsgpt" in parsed