mirror of
https://github.com/tiennm99/DocsGPT.git
synced 2026-10-11 03:12:55 +00:00
feat: compatible api
This commit is contained in:
1 parent
73256389cf
commit
addf57cab7
15 files changed
+1988
-8
No files matched your search
@@ -0,0 +1,3 @@
|
||||
from application.api.v1.routes import v1_bp
|
||||
|
||||
__all__ = ["v1_bp"]
|
||||
@@ -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,
|
||||
)
|
||||
@@ -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
|
||||
@@ -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,
|
||||
|
||||
@@ -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;
|
||||
};
|
||||
@@ -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)}
|
||||
|
||||
@@ -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'
|
||||
|
||||
@@ -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
|
||||
@@ -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
|
||||
Reference in new issue
Block a user