mirror of
https://github.com/open-webui/open-webui.git
synced 2026-10-07 02:58:21 +00:00
refac
This commit is contained in:
parent
0f46d6096c
commit
015dbc8619
5 changed files with 235 additions and 328 deletions
|
|
@ -248,9 +248,9 @@ from open_webui.utils.logger import start_logger
|
|||
from open_webui.utils.middleware import (
|
||||
background_tasks_handler,
|
||||
build_chat_response_context,
|
||||
drain_approved_tool_calls,
|
||||
process_chat_payload,
|
||||
process_chat_response,
|
||||
resume_tool_calls,
|
||||
)
|
||||
from open_webui.utils.misc import get_response_error_detail, merge_model_params
|
||||
from open_webui.utils.model_ids import strip_provider_model_prefix
|
||||
|
|
@ -1654,82 +1654,79 @@ async def chat_completion(
|
|||
)
|
||||
|
||||
async def process_chat(request, form_data, user, metadata, model, tasks=None):
|
||||
error_detail = None
|
||||
try:
|
||||
ctx = None
|
||||
# Saved chats load the message after approved tool calls run, so their results are kept
|
||||
if metadata.get('assistant_message_id') and not is_saved_chat_id(metadata.get('chat_id')):
|
||||
ctx = await build_chat_response_context(request, form_data, user, model, metadata, tasks, [])
|
||||
form_data, metadata, events = await process_chat_payload(request, form_data, user, metadata, model)
|
||||
|
||||
if await drain_approved_tool_calls(request, form_data, user, model, metadata):
|
||||
return {'status': True, 'chat_id': metadata.get('chat_id'), 'paused': True}
|
||||
|
||||
response = await chat_completion_handler(request, form_data, user)
|
||||
|
||||
# When the upstream provider returns an error (e.g. HTTP 400
|
||||
# content-filter, quota exceeded), generate_chat_completion
|
||||
# returns a JSONResponse instead of raising. Detect this and
|
||||
# raise so the except-block below emits a terminal
|
||||
# chat:message:error, unblocking the frontend.
|
||||
if isinstance(response, JSONResponse) and response.status_code >= 400:
|
||||
raise Exception(get_response_error_detail(response))
|
||||
|
||||
if ctx is None:
|
||||
ctx = await build_chat_response_context(request, form_data, user, model, metadata, tasks, events)
|
||||
else:
|
||||
ctx.update(form_data=form_data, metadata=metadata, events=events)
|
||||
|
||||
return await process_chat_response(response, ctx)
|
||||
except asyncio.CancelledError:
|
||||
log.info('Chat processing was cancelled')
|
||||
try:
|
||||
ctx = None
|
||||
# Saved chats load the message after approved tool calls run, so their results are kept
|
||||
if metadata.get('assistant_message_id') and not is_saved_chat_id(metadata.get('chat_id')):
|
||||
ctx = await build_chat_response_context(request, form_data, user, model, metadata, tasks, [])
|
||||
form_data, metadata, events = await process_chat_payload(request, form_data, user, metadata, model)
|
||||
|
||||
async def emit_cancel_event():
|
||||
event_emitter = await get_event_emitter(metadata)
|
||||
if event_emitter:
|
||||
await event_emitter({'type': 'chat:tasks:cancel'})
|
||||
paused = await resume_tool_calls(request, form_data, user, model, metadata)
|
||||
if paused:
|
||||
return {'status': True, 'chat_id': metadata.get('chat_id'), 'paused': True}
|
||||
|
||||
await asyncio.shield(emit_cancel_event())
|
||||
except Exception:
|
||||
pass
|
||||
raise # re-raise to ensure proper task cancellation handling
|
||||
except Exception as e:
|
||||
error_detail = e.detail if isinstance(e, HTTPException) else str(e)
|
||||
log.error('Error processing chat payload: %s', error_detail)
|
||||
if metadata.get('chat_id') and metadata.get('message_id'):
|
||||
# Update the chat message with the error
|
||||
response = await chat_completion_handler(request, form_data, user)
|
||||
|
||||
if isinstance(response, Response) and response.status_code >= 400:
|
||||
error_detail = get_response_error_detail(response)
|
||||
if metadata.get('session_id') and metadata.get('chat_id'):
|
||||
return None
|
||||
return response
|
||||
|
||||
if ctx is None:
|
||||
ctx = await build_chat_response_context(request, form_data, user, model, metadata, tasks, events)
|
||||
else:
|
||||
ctx.update(form_data=form_data, metadata=metadata, events=events)
|
||||
|
||||
return await process_chat_response(response, ctx)
|
||||
except asyncio.CancelledError:
|
||||
log.info('Chat processing was cancelled')
|
||||
try:
|
||||
if is_saved_chat_id(metadata.get('chat_id')):
|
||||
await Chats.upsert_message_to_chat_by_id_and_message_id(
|
||||
metadata['chat_id'],
|
||||
metadata['message_id'],
|
||||
{
|
||||
'parentId': metadata.get('user_message_id', None),
|
||||
'error': {'content': error_detail},
|
||||
'done': True,
|
||||
},
|
||||
)
|
||||
|
||||
event_emitter = await get_event_emitter(metadata)
|
||||
if event_emitter:
|
||||
await event_emitter(
|
||||
{
|
||||
'type': 'chat:message:error',
|
||||
'data': {'error': {'content': error_detail}, 'done': True},
|
||||
}
|
||||
)
|
||||
async def emit_cancel_event():
|
||||
event_emitter = await get_event_emitter(metadata)
|
||||
if event_emitter:
|
||||
await event_emitter({'type': 'chat:tasks:cancel'})
|
||||
|
||||
await asyncio.shield(emit_cancel_event())
|
||||
except Exception:
|
||||
pass
|
||||
else:
|
||||
# No chat_id/message_id → legacy/direct API path with no
|
||||
# WebSocket error channel. We must surface the error as
|
||||
# a proper HTTP response; without this the function would
|
||||
# return None which FastAPI serializes as null. #23924
|
||||
raise HTTPException(
|
||||
status_code=status.HTTP_400_BAD_REQUEST,
|
||||
detail=error_detail,
|
||||
)
|
||||
raise # re-raise to ensure proper task cancellation handling
|
||||
except Exception as e:
|
||||
error_detail = e.detail if isinstance(e, HTTPException) else str(e)
|
||||
if not (metadata.get('session_id') and metadata.get('chat_id')):
|
||||
raise
|
||||
finally:
|
||||
if error_detail is not None:
|
||||
log.error('Error processing chat payload: %s', error_detail)
|
||||
if metadata.get('chat_id') and metadata.get('message_id'):
|
||||
if is_saved_chat_id(metadata['chat_id']):
|
||||
try:
|
||||
await Chats.upsert_message_to_chat_by_id_and_message_id(
|
||||
metadata['chat_id'],
|
||||
metadata['message_id'],
|
||||
{
|
||||
'parentId': metadata.get('user_message_id'),
|
||||
'error': {'content': error_detail},
|
||||
'done': True,
|
||||
},
|
||||
)
|
||||
except Exception:
|
||||
log.exception('Failed to save chat error')
|
||||
|
||||
try:
|
||||
event_emitter = await get_event_emitter(metadata)
|
||||
if event_emitter:
|
||||
await event_emitter(
|
||||
{
|
||||
'type': 'chat:message:error',
|
||||
'data': {'error': {'content': error_detail}, 'done': True},
|
||||
}
|
||||
)
|
||||
except Exception:
|
||||
log.exception('Failed to emit chat error')
|
||||
finally:
|
||||
# Clean up MCP clients. Each client is isolated so one
|
||||
# failure doesn't skip the rest.
|
||||
|
|
@ -1895,7 +1892,12 @@ async def chat_completion(
|
|||
else:
|
||||
# Legacy/direct: single model, synchronous
|
||||
metadata['message_id'] = message_ids[0]['message_id']
|
||||
return await process_chat(request, form_data, user, metadata, model, tasks)
|
||||
try:
|
||||
return await process_chat(request, form_data, user, metadata, model, tasks)
|
||||
except HTTPException:
|
||||
raise
|
||||
except Exception as e:
|
||||
raise HTTPException(status_code=status.HTTP_400_BAD_REQUEST, detail=str(e)) from e
|
||||
|
||||
|
||||
# Alias for chat_completion (Legacy)
|
||||
|
|
@ -1997,9 +1999,12 @@ async def passthrough_anthropic_messages(request: Request, form_data: dict, user
|
|||
requested_model=requested_model,
|
||||
upstream_error=response_data,
|
||||
)
|
||||
retry_headers = {
|
||||
k: v for k, v in response.headers.items() if k.lower() in ('retry-after', 'retry-after-ms')
|
||||
}
|
||||
if isinstance(response_data, (dict, list)):
|
||||
return JSONResponse(status_code=response.status, content=response_data)
|
||||
return Response(status_code=response.status, content=response_data)
|
||||
return JSONResponse(status_code=response.status, content=response_data, headers=retry_headers)
|
||||
return Response(status_code=response.status, content=response_data, headers=retry_headers)
|
||||
|
||||
return response_data
|
||||
except HTTPException:
|
||||
|
|
|
|||
|
|
@ -122,6 +122,7 @@ async def send_request(
|
|||
)
|
||||
|
||||
if not r.ok:
|
||||
retry_headers = {k: v for k, v in r.headers.items() if k.lower() in ('retry-after', 'retry-after-ms')}
|
||||
try:
|
||||
res = await r.json(loads=JSONCodec.loads)
|
||||
await publish_model_provider_request_failed(
|
||||
|
|
@ -133,7 +134,7 @@ async def send_request(
|
|||
upstream_error=res,
|
||||
)
|
||||
if 'error' in res:
|
||||
raise HTTPException(status_code=r.status, detail=res['error'])
|
||||
raise HTTPException(status_code=r.status, detail=res['error'], headers=retry_headers)
|
||||
except HTTPException:
|
||||
raise
|
||||
except Exception as e:
|
||||
|
|
@ -148,6 +149,7 @@ async def send_request(
|
|||
raise HTTPException(
|
||||
status_code=r.status,
|
||||
detail=ERROR_MESSAGES.SERVER_CONNECTION_ERROR,
|
||||
headers=retry_headers,
|
||||
)
|
||||
|
||||
r.raise_for_status()
|
||||
|
|
|
|||
|
|
@ -1591,6 +1591,7 @@ async def generate_chat_completion(
|
|||
# read the body and return a proper error response instead of
|
||||
# streaming the error back (which hides the error from logs).
|
||||
if r.status >= 400:
|
||||
retry_headers = {k: v for k, v in r.headers.items() if k.lower() in ('retry-after', 'retry-after-ms')}
|
||||
error_body = await r.text()
|
||||
log.error(
|
||||
'Provider returned HTTP %d with SSE content-type: %s',
|
||||
|
|
@ -1609,7 +1610,7 @@ async def generate_chat_completion(
|
|||
requested_model=requested_model,
|
||||
upstream_error=error_json,
|
||||
)
|
||||
return JSONResponse(status_code=r.status, content=error_json)
|
||||
return JSONResponse(status_code=r.status, content=error_json, headers=retry_headers)
|
||||
except JSONCodec.JSONDecodeError:
|
||||
await publish_model_provider_request_failed(
|
||||
request,
|
||||
|
|
@ -1624,6 +1625,7 @@ async def generate_chat_completion(
|
|||
return JSONResponse(
|
||||
status_code=r.status,
|
||||
content={'error': {'message': error_body, 'code': r.status}},
|
||||
headers=retry_headers,
|
||||
)
|
||||
|
||||
streaming = True
|
||||
|
|
@ -1640,6 +1642,7 @@ async def generate_chat_completion(
|
|||
response = await r.text()
|
||||
|
||||
if r.status >= 400:
|
||||
retry_headers = {k: v for k, v in r.headers.items() if k.lower() in ('retry-after', 'retry-after-ms')}
|
||||
await publish_model_provider_request_failed(
|
||||
request,
|
||||
actor=user,
|
||||
|
|
@ -1651,9 +1654,9 @@ async def generate_chat_completion(
|
|||
upstream_error=response,
|
||||
)
|
||||
if isinstance(response, (dict, list)):
|
||||
return JSONResponse(status_code=r.status, content=response)
|
||||
return JSONResponse(status_code=r.status, content=response, headers=retry_headers)
|
||||
else:
|
||||
return PlainTextResponse(status_code=r.status, content=response)
|
||||
return PlainTextResponse(status_code=r.status, content=response, headers=retry_headers)
|
||||
|
||||
# Convert Responses API result to simple format
|
||||
if is_responses and isinstance(response, dict):
|
||||
|
|
@ -1751,6 +1754,7 @@ async def embeddings(request: Request, form_data: dict, user):
|
|||
response_data = await r.text()
|
||||
|
||||
if r.status >= 400:
|
||||
retry_headers = {k: v for k, v in r.headers.items() if k.lower() in ('retry-after', 'retry-after-ms')}
|
||||
await publish_model_provider_request_failed(
|
||||
request,
|
||||
actor=user,
|
||||
|
|
@ -1762,9 +1766,9 @@ async def embeddings(request: Request, form_data: dict, user):
|
|||
upstream_error=response_data,
|
||||
)
|
||||
if isinstance(response_data, (dict, list)):
|
||||
return JSONResponse(status_code=r.status, content=response_data)
|
||||
return JSONResponse(status_code=r.status, content=response_data, headers=retry_headers)
|
||||
else:
|
||||
return PlainTextResponse(status_code=r.status, content=response_data)
|
||||
return PlainTextResponse(status_code=r.status, content=response_data, headers=retry_headers)
|
||||
|
||||
return response_data
|
||||
except Exception as e:
|
||||
|
|
@ -1879,6 +1883,7 @@ async def responses(
|
|||
response_data = await r.text()
|
||||
|
||||
if r.status >= 400:
|
||||
retry_headers = {k: v for k, v in r.headers.items() if k.lower() in ('retry-after', 'retry-after-ms')}
|
||||
await publish_model_provider_request_failed(
|
||||
request,
|
||||
actor=user,
|
||||
|
|
@ -1890,9 +1895,9 @@ async def responses(
|
|||
upstream_error=response_data,
|
||||
)
|
||||
if isinstance(response_data, (dict, list)):
|
||||
return JSONResponse(status_code=r.status, content=response_data)
|
||||
return JSONResponse(status_code=r.status, content=response_data, headers=retry_headers)
|
||||
else:
|
||||
return PlainTextResponse(status_code=r.status, content=response_data)
|
||||
return PlainTextResponse(status_code=r.status, content=response_data, headers=retry_headers)
|
||||
|
||||
return response_data
|
||||
|
||||
|
|
@ -2001,6 +2006,7 @@ async def proxy(path: str, request: Request, user=Depends(get_verified_user)):
|
|||
response_data = await r.text()
|
||||
|
||||
if r.status >= 400:
|
||||
retry_headers = {k: v for k, v in r.headers.items() if k.lower() in ('retry-after', 'retry-after-ms')}
|
||||
await publish_model_provider_request_failed(
|
||||
request,
|
||||
actor=user,
|
||||
|
|
@ -2012,9 +2018,9 @@ async def proxy(path: str, request: Request, user=Depends(get_verified_user)):
|
|||
upstream_error=response_data,
|
||||
)
|
||||
if isinstance(response_data, (dict, list)):
|
||||
return JSONResponse(status_code=r.status, content=response_data)
|
||||
return JSONResponse(status_code=r.status, content=response_data, headers=retry_headers)
|
||||
else:
|
||||
return PlainTextResponse(status_code=r.status, content=response_data)
|
||||
return PlainTextResponse(status_code=r.status, content=response_data, headers=retry_headers)
|
||||
|
||||
return response_data
|
||||
|
||||
|
|
|
|||
|
|
@ -3295,8 +3295,7 @@ async def build_chat_response_context(request, form_data, user, model, metadata,
|
|||
}
|
||||
|
||||
|
||||
async def execute_tool_call_for_output(request, form_data, user, metadata, event_caller, event_emitter, tool_call):
|
||||
tools = metadata.get('tools', {})
|
||||
async def execute_tool_call(form_data, metadata, event_caller, tool_call):
|
||||
name = tool_call.get('function', {}).get('name', '')
|
||||
tool_args = tool_call.get('function', {}).get('arguments', '{}')
|
||||
params = {}
|
||||
|
|
@ -3308,28 +3307,23 @@ async def execute_tool_call_for_output(request, form_data, user, metadata, event
|
|||
params = ast.literal_eval(tool_args)
|
||||
except Exception as e:
|
||||
log.debug(e)
|
||||
return {
|
||||
'tool_call_id': tool_call.get('id', ''),
|
||||
'content': (
|
||||
'Error: Tool call arguments could not be parsed. '
|
||||
'The model generated malformed or incomplete JSON.'
|
||||
),
|
||||
}
|
||||
return {}, None, None, None, False
|
||||
if not isinstance(params, dict):
|
||||
return {
|
||||
'tool_call_id': tool_call.get('id', ''),
|
||||
'content': 'Error: Tool call arguments must be a JSON object.',
|
||||
}
|
||||
return (
|
||||
{},
|
||||
f'Error: Tool call arguments for `{name}` must be a JSON object. Please try again.',
|
||||
None,
|
||||
None,
|
||||
False,
|
||||
)
|
||||
tool_call.setdefault('function', {})['arguments'] = JSONCodec.dumps(params)
|
||||
|
||||
tool = tools.get(name)
|
||||
tool = metadata.get('tools', {}).get(name)
|
||||
if not tool:
|
||||
return {'tool_call_id': tool_call.get('id', ''), 'content': f'Error: Tool "{name}" not found.'}
|
||||
|
||||
spec = tool.get('spec', {})
|
||||
return params, f'Error: Tool "{name}" not found.', None, None, False
|
||||
tool_type = tool.get('type', '')
|
||||
direct_tool = tool.get('direct', False)
|
||||
allowed_params = spec.get('parameters', {}).get('properties', {}).keys()
|
||||
allowed_params = tool.get('spec', {}).get('parameters', {}).get('properties', {}).keys()
|
||||
params = {key: value for key, value in params.items() if key in allowed_params}
|
||||
|
||||
try:
|
||||
|
|
@ -3360,35 +3354,13 @@ async def execute_tool_call_for_output(request, form_data, user, metadata, event
|
|||
result = await function(**params)
|
||||
except Exception as e:
|
||||
result = {'error': str(e)}
|
||||
|
||||
terminal_file_result = build_terminal_file_tool_result(name, params, result, tool, metadata)
|
||||
if terminal_file_result:
|
||||
result = terminal_file_result
|
||||
|
||||
result, files, embeds = await process_tool_result(
|
||||
request,
|
||||
name,
|
||||
result,
|
||||
tool_type,
|
||||
direct_tool,
|
||||
metadata,
|
||||
user,
|
||||
)
|
||||
|
||||
await terminal_event_handler(name, params, result, event_emitter)
|
||||
|
||||
return {
|
||||
'tool_call_id': tool_call.get('id', ''),
|
||||
'content': tool_result_content(result),
|
||||
**({'files': files} if files else {}),
|
||||
**({'embeds': embeds} if embeds else {}),
|
||||
}
|
||||
return params, result, tool, tool_type, direct_tool
|
||||
|
||||
|
||||
async def drain_approved_tool_calls(request, form_data, user, model, metadata) -> bool:
|
||||
async def resume_tool_calls(request, form_data, user, model, metadata) -> bool:
|
||||
"""Execute approved calls on a saved message; return whether it is still paused."""
|
||||
chat_id = metadata.get('chat_id')
|
||||
assistant_message_id = metadata.get('assistant_message_id')
|
||||
# Only a resume/continue payload re-enters an existing message; other paths mint a fresh id with nothing to drain.
|
||||
if not is_saved_chat_id(chat_id) or not assistant_message_id:
|
||||
return False
|
||||
|
||||
|
|
@ -3410,30 +3382,23 @@ async def drain_approved_tool_calls(request, form_data, user, model, metadata) -
|
|||
and item.get('approved') is True
|
||||
and item.get('call_id') not in result_call_ids
|
||||
]
|
||||
if not approved_calls:
|
||||
if metadata.get('params', {}).get('tool_approval_mode', 'full') == 'ask' and any(
|
||||
item.get('type') == 'function_call'
|
||||
and item.get('name') != 'ask_user'
|
||||
and (item.get('call_id') or item.get('id'))
|
||||
and item.get('status') == 'queued'
|
||||
and item.get('approved') is not True
|
||||
and (item.get('call_id') or item.get('id')) not in result_call_ids
|
||||
for item in output
|
||||
):
|
||||
event_emitter, _ = await get_event_emitter_and_caller(metadata)
|
||||
await pause_for_tool_approval(chat_id, message_id, output, form_data, metadata)
|
||||
if event_emitter:
|
||||
await event_emitter({'type': 'chat:completion', 'data': {'done': False, 'output': output}})
|
||||
return True
|
||||
needs_approval = metadata.get('params', {}).get('tool_approval_mode', 'full') == 'ask' and any(
|
||||
item.get('type') == 'function_call'
|
||||
and item.get('name') != 'ask_user'
|
||||
and (item.get('call_id') or item.get('id'))
|
||||
and item.get('status') == 'queued'
|
||||
and item.get('approved') is not True
|
||||
and (item.get('call_id') or item.get('id')) not in result_call_ids
|
||||
for item in output
|
||||
)
|
||||
if not approved_calls and not needs_approval:
|
||||
return False
|
||||
|
||||
event_emitter, event_caller = await get_event_emitter_and_caller(metadata)
|
||||
changed = False
|
||||
for item in approved_calls:
|
||||
if item.get('name') == 'ask_user':
|
||||
item['status'] = 'pending'
|
||||
item.pop('approved', None)
|
||||
changed = True
|
||||
continue
|
||||
|
||||
tool_call = {
|
||||
|
|
@ -3444,20 +3409,27 @@ async def drain_approved_tool_calls(request, form_data, user, model, metadata) -
|
|||
'arguments': item.get('arguments', '{}'),
|
||||
},
|
||||
}
|
||||
result = await execute_tool_call_for_output(
|
||||
request,
|
||||
form_data,
|
||||
user,
|
||||
metadata,
|
||||
event_caller,
|
||||
event_emitter,
|
||||
tool_call,
|
||||
params, result, tool, tool_type, direct_tool = await execute_tool_call(
|
||||
form_data, metadata, event_caller, tool_call
|
||||
)
|
||||
files, embeds = [], []
|
||||
if result is None and tool is None:
|
||||
result = 'Error: Tool call arguments could not be parsed. The model generated malformed or incomplete JSON.'
|
||||
elif tool:
|
||||
name = item.get('name', '')
|
||||
terminal_file_result = build_terminal_file_tool_result(name, params, result, tool, metadata)
|
||||
if terminal_file_result:
|
||||
result = terminal_file_result
|
||||
result, files, embeds = await process_tool_result(
|
||||
request, name, result, tool_type, direct_tool, metadata, user
|
||||
)
|
||||
await terminal_event_handler(name, params, result, event_emitter)
|
||||
content = tool_result_content(result)
|
||||
item['arguments'] = tool_call.get('function', {}).get('arguments', '{}')
|
||||
output_parts = [{'type': 'input_text', 'text': result.get('content', '')}]
|
||||
item['status'] = 'failed' if _is_tool_result_error(result.get('content', '')) else 'completed'
|
||||
output_parts = [{'type': 'input_text', 'text': content}]
|
||||
item['status'] = 'failed' if _is_tool_result_error(content) else 'completed'
|
||||
display_files = []
|
||||
for file_item in result.get('files', []):
|
||||
for file_item in files:
|
||||
if file_item.get('type') == 'image' and file_item.get('url', '').startswith('data:'):
|
||||
image_url = await store_tool_result_image(request, file_item['url'], metadata, user)
|
||||
output_parts.append({'type': 'input_image', 'image_url': image_url})
|
||||
|
|
@ -3470,128 +3442,107 @@ async def drain_approved_tool_calls(request, form_data, user, model, metadata) -
|
|||
{
|
||||
'type': 'function_call_output',
|
||||
'id': output_id('fco'),
|
||||
'call_id': result.get('tool_call_id', ''),
|
||||
'call_id': tool_call['id'],
|
||||
'output': output_parts,
|
||||
'status': item['status'],
|
||||
**({'files': display_files} if display_files else {}),
|
||||
**({'embeds': result.get('embeds')} if result.get('embeds') else {}),
|
||||
**({'embeds': embeds} if embeds else {}),
|
||||
}
|
||||
)
|
||||
changed = True
|
||||
result_call_ids.add(tool_call['id'])
|
||||
|
||||
if changed:
|
||||
result_call_ids = {
|
||||
item.get('call_id') for item in output if item.get('type') == 'function_call_output' and item.get('call_id')
|
||||
}
|
||||
if metadata.get('params', {}).get('tool_approval_mode', 'full') == 'ask' and any(
|
||||
item.get('type') == 'function_call'
|
||||
and item.get('name') != 'ask_user'
|
||||
and (item.get('call_id') or item.get('id'))
|
||||
and item.get('status') == 'queued'
|
||||
and item.get('approved') is not True
|
||||
and (item.get('call_id') or item.get('id')) not in result_call_ids
|
||||
for item in output
|
||||
):
|
||||
await pause_for_tool_approval(chat_id, message_id, output, form_data, metadata)
|
||||
result_call_ids = {
|
||||
item.get('call_id')
|
||||
for item in output
|
||||
if item.get('type') == 'function_call_output' and item.get('call_id')
|
||||
if needs_approval:
|
||||
await pause_for_tool_approval(chat_id, message_id, output, form_data, metadata)
|
||||
paused = any(
|
||||
item.get('type') == 'function_call'
|
||||
and item.get('call_id')
|
||||
and item.get('status') in {'pending', 'queued', 'requires_approval'}
|
||||
and item.get('call_id') not in result_call_ids
|
||||
for item in output
|
||||
)
|
||||
if not paused:
|
||||
output.append(
|
||||
{
|
||||
'type': 'message',
|
||||
'id': output_id('msg'),
|
||||
'status': 'in_progress',
|
||||
'role': 'assistant',
|
||||
'content': [{'type': 'output_text', 'text': ''}],
|
||||
}
|
||||
paused = any(
|
||||
item.get('type') == 'function_call'
|
||||
and item.get('call_id')
|
||||
and item.get('status') in {'pending', 'queued', 'requires_approval'}
|
||||
and item.get('call_id') not in result_call_ids
|
||||
for item in output
|
||||
)
|
||||
if not paused:
|
||||
output.append(
|
||||
{
|
||||
'type': 'message',
|
||||
'id': output_id('msg'),
|
||||
'status': 'in_progress',
|
||||
'role': 'assistant',
|
||||
'content': [{'type': 'output_text', 'text': ''}],
|
||||
}
|
||||
)
|
||||
|
||||
if not needs_approval:
|
||||
await Chats.upsert_message_to_chat_by_id_and_message_id(
|
||||
chat_id,
|
||||
message_id,
|
||||
{'done': False, 'output': output},
|
||||
touch=False,
|
||||
)
|
||||
if event_emitter:
|
||||
await event_emitter(
|
||||
{
|
||||
'type': 'chat:completion',
|
||||
'data': {
|
||||
'done': False,
|
||||
'output': output,
|
||||
},
|
||||
}
|
||||
if event_emitter:
|
||||
await event_emitter(
|
||||
{
|
||||
'type': 'chat:completion',
|
||||
'data': {
|
||||
'done': False,
|
||||
'output': output,
|
||||
},
|
||||
}
|
||||
)
|
||||
|
||||
if paused:
|
||||
return True
|
||||
|
||||
db_messages = await load_messages_from_db(chat_id, metadata.get('user_message_id'))
|
||||
if db_messages:
|
||||
assistant_message = await Chats.get_message_by_id_and_message_id(chat_id, message_id)
|
||||
if assistant_message:
|
||||
db_messages.append({k: v for k, v in assistant_message.items() if k in MESSAGE_REPLAY_KEYS})
|
||||
context_start_message_id = metadata.get('context_start_message_id')
|
||||
start_index = next(
|
||||
(index for index, message in enumerate(db_messages) if message.get('id') == context_start_message_id), 0
|
||||
)
|
||||
db_messages = db_messages[start_index:]
|
||||
for message in db_messages:
|
||||
output = message.get('output')
|
||||
# reasoning_details can be model/provider-bound, so only replay them
|
||||
# for output produced by the same model.
|
||||
if message.get('role') == 'assistant' and message.get('model') != model['id'] and isinstance(output, list):
|
||||
message['output'] = strip_reasoning_details(output)
|
||||
|
||||
system_message = get_system_message(form_data.get('messages', []))
|
||||
form_data['messages'] = process_messages_with_output(
|
||||
[system_message, *db_messages] if system_message else db_messages,
|
||||
reasoning_format=get_reasoning_format(model),
|
||||
include_file_context=metadata.get('include_file_context', False),
|
||||
)
|
||||
form_data['messages'] = sanitize_tool_pairs(form_data['messages'])
|
||||
|
||||
if ENABLE_FUNCTIONS:
|
||||
filter_functions = await get_filter_functions(request, model, metadata.get('filter_ids', []))
|
||||
if filter_functions:
|
||||
filtered_form_data, _ = await process_filter_functions(
|
||||
request=request,
|
||||
filter_context=get_filter_context(request),
|
||||
filter_functions=filter_functions,
|
||||
filter_type='request',
|
||||
form_data=form_data,
|
||||
extra_params={
|
||||
'__event_emitter__': event_emitter,
|
||||
'__event_call__': event_caller,
|
||||
'__user__': user.model_dump() if isinstance(user, UserModel) else {},
|
||||
'__metadata__': metadata,
|
||||
'__oauth_token__': await get_system_oauth_token(request, user),
|
||||
'__request__': request,
|
||||
'__model__': model,
|
||||
'__chat_id__': metadata.get('chat_id'),
|
||||
'__message_id__': metadata.get('message_id'),
|
||||
},
|
||||
)
|
||||
if filtered_form_data is not form_data:
|
||||
form_data.clear()
|
||||
form_data.update(filtered_form_data)
|
||||
|
||||
db_messages = await load_messages_from_db(chat_id, metadata.get('user_message_id'))
|
||||
if db_messages:
|
||||
assistant_message = await Chats.get_message_by_id_and_message_id(chat_id, message_id)
|
||||
if assistant_message:
|
||||
db_messages.append({k: v for k, v in assistant_message.items() if k in MESSAGE_REPLAY_KEYS})
|
||||
context_start_message_id = metadata.get('context_start_message_id')
|
||||
start_index = next(
|
||||
(index for index, message in enumerate(db_messages) if message.get('id') == context_start_message_id), 0
|
||||
)
|
||||
db_messages = db_messages[start_index:]
|
||||
for message in db_messages:
|
||||
output = message.get('output')
|
||||
# reasoning_details can be model/provider-bound, so only replay them
|
||||
# for output produced by the same model.
|
||||
if (
|
||||
message.get('role') == 'assistant'
|
||||
and message.get('model') != model['id']
|
||||
and isinstance(output, list)
|
||||
):
|
||||
message['output'] = strip_reasoning_details(output)
|
||||
|
||||
system_message = get_system_message(form_data.get('messages', []))
|
||||
form_data['messages'] = process_messages_with_output(
|
||||
[system_message, *db_messages] if system_message else db_messages,
|
||||
reasoning_format=get_reasoning_format(model),
|
||||
include_file_context=metadata.get('include_file_context', False),
|
||||
)
|
||||
form_data['messages'] = sanitize_tool_pairs(form_data['messages'])
|
||||
|
||||
if not paused and ENABLE_FUNCTIONS:
|
||||
filter_functions = await get_filter_functions(request, model, metadata.get('filter_ids', []))
|
||||
if filter_functions:
|
||||
filtered_form_data, _ = await process_filter_functions(
|
||||
request=request,
|
||||
filter_context=get_filter_context(request),
|
||||
filter_functions=filter_functions,
|
||||
filter_type='request',
|
||||
form_data=form_data,
|
||||
extra_params={
|
||||
'__event_emitter__': event_emitter,
|
||||
'__event_call__': event_caller,
|
||||
'__user__': user.model_dump() if isinstance(user, UserModel) else {},
|
||||
'__metadata__': metadata,
|
||||
'__oauth_token__': await get_system_oauth_token(request, user),
|
||||
'__request__': request,
|
||||
'__model__': model,
|
||||
'__chat_id__': metadata.get('chat_id'),
|
||||
'__message_id__': metadata.get('message_id'),
|
||||
},
|
||||
)
|
||||
if filtered_form_data is not form_data:
|
||||
form_data.clear()
|
||||
form_data.update(filtered_form_data)
|
||||
|
||||
if not paused:
|
||||
normalize_messages_for_model(form_data)
|
||||
|
||||
return paused
|
||||
|
||||
normalize_messages_for_model(form_data)
|
||||
return False
|
||||
|
||||
|
||||
|
|
@ -5959,76 +5910,8 @@ async def streaming_chat_response_handler(response, ctx):
|
|||
|
||||
await emit_output()
|
||||
|
||||
tools = metadata.get('tools', {})
|
||||
|
||||
results = []
|
||||
|
||||
def parse_tool_params(tool_call):
|
||||
tool_args = tool_call.get('function', {}).get('arguments', '{}')
|
||||
params = {}
|
||||
if tool_args and tool_args.strip():
|
||||
try:
|
||||
params = JSONCodec.loads(tool_args)
|
||||
except Exception:
|
||||
try:
|
||||
params = ast.literal_eval(tool_args)
|
||||
except Exception as e:
|
||||
log.debug(e)
|
||||
return None
|
||||
if not isinstance(params, dict):
|
||||
raise ValueError('Tool call arguments must be a JSON object.')
|
||||
tool_call.setdefault('function', {})['arguments'] = JSONCodec.dumps(params)
|
||||
return params
|
||||
|
||||
async def execute_tool_call(tool_call):
|
||||
name = tool_call.get('function', {}).get('name', '')
|
||||
try:
|
||||
params = parse_tool_params(tool_call)
|
||||
except ValueError:
|
||||
return (
|
||||
{},
|
||||
f'Error: Tool call arguments for `{name}` must be a JSON object. Please try again.',
|
||||
None,
|
||||
None,
|
||||
False,
|
||||
)
|
||||
if params is None:
|
||||
return {}, None, None, None, False
|
||||
tool = tools.get(name)
|
||||
if not tool:
|
||||
return params, f'Error: Tool "{name}" not found.', None, None, False
|
||||
spec = tool.get('spec', {})
|
||||
tool_type = tool.get('type', '')
|
||||
direct_tool = tool.get('direct', False)
|
||||
allowed_params = spec.get('parameters', {}).get('properties', {}).keys()
|
||||
params = {key: value for key, value in params.items() if key in allowed_params}
|
||||
try:
|
||||
if direct_tool:
|
||||
result = await event_caller(
|
||||
{
|
||||
'type': 'execute:tool',
|
||||
'data': {
|
||||
'id': str(uuid4()),
|
||||
'name': name,
|
||||
'params': params,
|
||||
'server': tool.get('server', {}),
|
||||
'session_id': metadata.get('session_id'),
|
||||
},
|
||||
}
|
||||
)
|
||||
else:
|
||||
function = await get_updated_tool_function(
|
||||
function=tool['callable'],
|
||||
extra_params={
|
||||
'__messages__': form_data.get('messages', []),
|
||||
'__files__': metadata.get('files', []),
|
||||
},
|
||||
)
|
||||
result = await function(**params)
|
||||
except Exception as e:
|
||||
result = {'error': str(e)}
|
||||
return params, result, tool, tool_type, direct_tool
|
||||
|
||||
delegate_calls = [
|
||||
tool_call
|
||||
for tool_call in response_tool_calls
|
||||
|
|
@ -6037,11 +5920,18 @@ async def streaming_chat_response_handler(response, ctx):
|
|||
tool_results = {}
|
||||
for tool_call in response_tool_calls:
|
||||
if tool_call.get('function', {}).get('name') != 'delegate_task':
|
||||
tool_results[id(tool_call)] = await execute_tool_call(tool_call)
|
||||
tool_results[id(tool_call)] = await execute_tool_call(
|
||||
form_data, metadata, event_caller, tool_call
|
||||
)
|
||||
tool_results.update(
|
||||
zip(
|
||||
[id(tool_call) for tool_call in delegate_calls],
|
||||
await asyncio.gather(*(execute_tool_call(tool_call) for tool_call in delegate_calls)),
|
||||
await asyncio.gather(
|
||||
*(
|
||||
execute_tool_call(form_data, metadata, event_caller, tool_call)
|
||||
for tool_call in delegate_calls
|
||||
)
|
||||
),
|
||||
)
|
||||
)
|
||||
|
||||
|
|
|
|||
|
|
@ -48,10 +48,14 @@ def get_response_error_detail(response: object) -> str:
|
|||
body = response.body
|
||||
if not isinstance(body, str):
|
||||
body = body.decode('utf-8', 'replace')
|
||||
detail = JSONCodec.loads(body)
|
||||
except Exception:
|
||||
return fallback
|
||||
|
||||
try:
|
||||
detail = JSONCodec.loads(body)
|
||||
except JSONCodec.JSONDecodeError:
|
||||
return body.strip() or fallback
|
||||
|
||||
while isinstance(detail, dict):
|
||||
next_detail = None
|
||||
for key in ('error', 'message', 'detail'):
|
||||
|
|
|
|||
Loading…
Add table
Reference in a new issue