mirror of
https://github.com/open-webui/open-webui.git
synced 2026-10-11 03:38:02 +00:00
refac
This commit is contained in:
parent
46a2a830ab
commit
93723bd315
4 changed files with 282 additions and 197 deletions
|
|
@ -1308,7 +1308,11 @@ async def chat_completion(
|
|||
or 'full'
|
||||
)
|
||||
|
||||
approval_resume = getattr(request.state, 'tool_approval_resume', None)
|
||||
metadata = {
|
||||
'tool_approval_resume': approval_resume[2]
|
||||
if approval_resume and approval_resume[:2] == (chat_id, form_data.get('assistant_message_id'))
|
||||
else None,
|
||||
'user_id': user.id,
|
||||
'user_agent': request.headers.get('user-agent', '') or '',
|
||||
'internal': getattr(request.state, 'internal', False) is True,
|
||||
|
|
@ -1916,8 +1920,14 @@ async def resolve_chat_message_tool_call(
|
|||
db: AsyncSession = Depends(get_async_session),
|
||||
):
|
||||
resolution = await resolve_tool_call_output(id, message_id, form_data, user, db=db)
|
||||
payload = await build_tool_approval_resume_payload(id, message_id, chat=resolution['chat'])
|
||||
result = await chat_completion(request, payload, user)
|
||||
if resolution['resume'] is None:
|
||||
return {'status': True, 'chat_id': id, 'message_id': message_id, 'task_ids': []}
|
||||
request.state.tool_approval_resume = (id, message_id, resolution['resume'])
|
||||
try:
|
||||
payload = await build_tool_approval_resume_payload(id, message_id, chat=resolution['chat'])
|
||||
result = await chat_completion(request, payload, user)
|
||||
finally:
|
||||
del request.state.tool_approval_resume
|
||||
return {
|
||||
'status': True,
|
||||
'chat_id': id,
|
||||
|
|
|
|||
|
|
@ -3,6 +3,7 @@
|
|||
from __future__ import annotations
|
||||
|
||||
from contextlib import asynccontextmanager
|
||||
from copy import deepcopy
|
||||
import logging
|
||||
import re
|
||||
import time
|
||||
|
|
@ -405,6 +406,38 @@ class ChatTable:
|
|||
yield session, chat
|
||||
await session.commit()
|
||||
|
||||
@asynccontextmanager
|
||||
async def edit_message_output(self, chat_id: str, message_id: str):
|
||||
"""Read and modify tool state under the same lock, updating both message stores."""
|
||||
async with self._chat_transaction(chat_id) as (session, chat):
|
||||
if chat is None:
|
||||
yield None
|
||||
return
|
||||
history = (chat.chat or {}).get('history') or {}
|
||||
message = dict(history.get('messages', {}).get(message_id) or {})
|
||||
row = await session.get(ChatMessage, f'{chat_id}-{message_id}')
|
||||
if row is not None:
|
||||
message.update(
|
||||
{
|
||||
ChatMessages.DB_TO_JSON_KEY_MAP.get(column.key, column.key): getattr(row, column.key)
|
||||
for column in ChatMessage.__table__.columns
|
||||
if column.key not in ChatMessages.EXCLUDED_COLUMNS
|
||||
}
|
||||
)
|
||||
message['id'] = message_id
|
||||
if not message:
|
||||
yield None
|
||||
return
|
||||
message = deepcopy(message)
|
||||
yield message
|
||||
message = self._clean_null_bytes(message)
|
||||
history.setdefault('messages', {})[message_id] = message
|
||||
chat.chat = {**(chat.chat or {}), 'history': history}
|
||||
flag_modified(chat, 'chat')
|
||||
await ChatMessages.upsert_message(
|
||||
message_id, chat_id, message.get('user_id') or chat.user_id, message, db=session
|
||||
)
|
||||
|
||||
def _clean_null_bytes(self, obj):
|
||||
"""Recursively remove null bytes from strings in dict/list structures."""
|
||||
return sanitize_data_for_db(obj)
|
||||
|
|
|
|||
|
|
@ -143,6 +143,7 @@ from open_webui.utils.task import (
|
|||
rag_template,
|
||||
tools_function_calling_generation_template,
|
||||
)
|
||||
from open_webui.utils.tool_approval import complete_tool_call, pending_tool_calls
|
||||
from open_webui.utils.tools import (
|
||||
connect_mcp_server,
|
||||
get_attached_knowledge,
|
||||
|
|
@ -2465,19 +2466,22 @@ async def process_chat_payload(request, form_data, user, metadata, model):
|
|||
if assistant_message_id:
|
||||
assistant_message = await Chats.get_message_by_id_and_message_id(chat_id, assistant_message_id)
|
||||
if assistant_message and (assistant_message.get('content') or assistant_message.get('output')):
|
||||
output = assistant_message.get('output') or []
|
||||
if (
|
||||
metadata.get('tool_approval_resume') != 'model'
|
||||
and not assistant_message.get('done', True)
|
||||
and output
|
||||
and output[-1].get('type') == 'message'
|
||||
and output[-1].get('status') == 'in_progress'
|
||||
and any(item.get('type') == 'function_call_output' for item in output)
|
||||
and not pending_tool_calls(output)
|
||||
):
|
||||
return form_data, metadata, [], True
|
||||
assistant_message = {k: v for k, v in assistant_message.items() if k in MESSAGE_REPLAY_KEYS}
|
||||
db_messages.append(assistant_message)
|
||||
output = assistant_message.get('output')
|
||||
output = output if isinstance(output, list) else []
|
||||
result_call_ids = {
|
||||
item.get('call_id') for item in output if item.get('type') == 'function_call_output'
|
||||
}
|
||||
if any(
|
||||
item.get('type') == 'function_call'
|
||||
and item.get('status') in {'pending', 'queued', 'requires_approval'}
|
||||
and item.get('call_id') not in result_call_ids
|
||||
for item in output
|
||||
):
|
||||
if metadata.get('tool_approval_resume') == 'tools' or pending_tool_calls(output):
|
||||
pending_assistant_message = assistant_message
|
||||
|
||||
system_message = get_system_message(form_data.get('messages', []))
|
||||
|
|
@ -3507,36 +3511,6 @@ async def resume_tool_calls(request, form_data, user, model, metadata, message)
|
|||
if not isinstance(output, list):
|
||||
return False
|
||||
|
||||
result_call_ids = {
|
||||
item.get('call_id') for item in output if item.get('type') == 'function_call_output' and item.get('call_id')
|
||||
}
|
||||
approved_calls = [
|
||||
item
|
||||
for item in output
|
||||
if item.get('type') == 'function_call'
|
||||
and item.get('call_id')
|
||||
and item.get('status') == 'queued'
|
||||
and item.get('approved') is True
|
||||
and item.get('call_id') not in result_call_ids
|
||||
]
|
||||
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 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
|
||||
)
|
||||
|
||||
event_emitter, event_caller = await get_event_emitter_and_caller(metadata)
|
||||
|
||||
# Request filters must be able to reject the complete request before tool side effects.
|
||||
|
|
@ -3577,11 +3551,34 @@ async def resume_tool_calls(request, form_data, user, model, metadata, message)
|
|||
)
|
||||
normalize_messages_for_model(form_data)
|
||||
|
||||
for item in approved_calls:
|
||||
if item.get('name') == 'ask_user':
|
||||
item['status'] = 'pending'
|
||||
item.pop('approved', None)
|
||||
continue
|
||||
while True:
|
||||
async with Chats.edit_message_output(chat_id, message_id) as saved:
|
||||
if not saved:
|
||||
raise HTTPException(status_code=404, detail='Message not found.')
|
||||
output = saved.get('output') or []
|
||||
pending = pending_tool_calls(output)
|
||||
if any(entry.get('status') == 'in_progress' for entry in pending):
|
||||
return True
|
||||
item = next(
|
||||
(
|
||||
entry
|
||||
for entry in pending
|
||||
if entry.get('name') != 'ask_user'
|
||||
and (
|
||||
entry.get('approved') is True or metadata.get('params', {}).get('tool_approval_mode') == 'full'
|
||||
)
|
||||
),
|
||||
None,
|
||||
)
|
||||
if item is None:
|
||||
if pending:
|
||||
pending[0]['status'] = 'pending'
|
||||
# A stale continuation must not start another model response.
|
||||
return True
|
||||
item['call_id'] = item.get('call_id') or item['id']
|
||||
item['status'] = 'in_progress'
|
||||
saved['done'] = False
|
||||
item = copy.deepcopy(item)
|
||||
|
||||
tool_call = {
|
||||
'id': item.get('call_id', ''),
|
||||
|
|
@ -3591,37 +3588,39 @@ async def resume_tool_calls(request, form_data, user, model, metadata, message)
|
|||
'arguments': item.get('arguments', '{}'),
|
||||
},
|
||||
}
|
||||
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
|
||||
try:
|
||||
params, result, tool, tool_type, direct_tool = await execute_tool_call(
|
||||
form_data, metadata, event_caller, tool_call
|
||||
)
|
||||
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': content}]
|
||||
item['status'] = 'failed' if _is_tool_result_error(content) else 'completed'
|
||||
display_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})
|
||||
else:
|
||||
display_files.append(file_item)
|
||||
if file_item.get('type') == 'image' and file_item.get('url'):
|
||||
output_parts.append({'type': 'input_image', 'image_url': file_item['url']})
|
||||
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': content}]
|
||||
item['status'] = 'failed' if _is_tool_result_error(content) else 'completed'
|
||||
display_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})
|
||||
else:
|
||||
display_files.append(file_item)
|
||||
if file_item.get('type') == 'image' and file_item.get('url'):
|
||||
output_parts.append({'type': 'input_image', 'image_url': file_item['url']})
|
||||
|
||||
output.append(
|
||||
{
|
||||
result_item = {
|
||||
'type': 'function_call_output',
|
||||
'id': output_id('fco'),
|
||||
'call_id': tool_call['id'],
|
||||
|
|
@ -3630,48 +3629,39 @@ async def resume_tool_calls(request, form_data, user, model, metadata, message)
|
|||
**({'files': display_files} if display_files else {}),
|
||||
**({'embeds': embeds} if embeds else {}),
|
||||
}
|
||||
)
|
||||
result_call_ids.add(tool_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': ''}],
|
||||
}
|
||||
)
|
||||
|
||||
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,
|
||||
},
|
||||
}
|
||||
)
|
||||
|
||||
return paused
|
||||
output = await complete_tool_call(chat_id, message_id, item, result_item)
|
||||
except BaseException:
|
||||
item['status'] = 'incomplete'
|
||||
await asyncio.shield(
|
||||
complete_tool_call(
|
||||
chat_id,
|
||||
message_id,
|
||||
item,
|
||||
{
|
||||
'type': 'function_call_output',
|
||||
'id': output_id('fco'),
|
||||
'call_id': tool_call['id'],
|
||||
'status': 'incomplete',
|
||||
'output': [
|
||||
{
|
||||
'type': 'input_text',
|
||||
'text': 'Error: tool execution was interrupted; its outcome is unknown.',
|
||||
}
|
||||
],
|
||||
},
|
||||
)
|
||||
)
|
||||
raise
|
||||
if event_emitter:
|
||||
await event_emitter({'type': 'chat:completion', 'data': {'done': False, 'output': output}})
|
||||
pending = pending_tool_calls(output)
|
||||
if not pending:
|
||||
message['output'] = output
|
||||
return False
|
||||
if metadata.get('params', {}).get('tool_approval_mode') != 'full' and not any(
|
||||
entry.get('approved') is True for entry in pending
|
||||
):
|
||||
return True
|
||||
|
||||
|
||||
async def pause_for_tool_approval(chat_id: str, message_id: str, output: list[dict], form_data: dict, metadata: dict):
|
||||
|
|
|
|||
|
|
@ -1,4 +1,5 @@
|
|||
from typing import Any, Literal
|
||||
from uuid import uuid4
|
||||
|
||||
from fastapi import HTTPException, status
|
||||
from pydantic import BaseModel
|
||||
|
|
@ -19,6 +20,50 @@ class ResolveToolCallForm(BaseModel):
|
|||
timed_out: bool = False
|
||||
|
||||
|
||||
def pending_tool_calls(output):
|
||||
result_ids = {item.get('call_id') for item in output if item.get('type') == 'function_call_output'}
|
||||
return [
|
||||
item
|
||||
for item in output
|
||||
if item.get('type') == 'function_call'
|
||||
and (item.get('call_id') or item.get('id')) not in result_ids
|
||||
and item.get('status') in {'pending', 'queued', 'requires_approval', 'in_progress'}
|
||||
]
|
||||
|
||||
|
||||
def tool_continuation():
|
||||
return {
|
||||
'type': 'message',
|
||||
'id': f'msg_{uuid4().hex}',
|
||||
'status': 'in_progress',
|
||||
'role': 'assistant',
|
||||
'content': [{'type': 'output_text', 'text': ''}],
|
||||
}
|
||||
|
||||
|
||||
async def complete_tool_call(chat_id, message_id, call, result):
|
||||
async with Chats.edit_message_output(chat_id, message_id) as message:
|
||||
if not message:
|
||||
raise HTTPException(status_code=404, detail=ERROR_MESSAGES.NOT_FOUND)
|
||||
output = message.get('output') or []
|
||||
saved = next((item for item in pending_tool_calls(output) if item.get('call_id') == call['call_id']), None)
|
||||
if saved is None or saved.get('status') != 'in_progress':
|
||||
raise HTTPException(status_code=409, detail='Tool call is no longer running.')
|
||||
saved.update(call)
|
||||
output.append(result)
|
||||
pending = pending_tool_calls(output)
|
||||
if result.get('status') == 'incomplete':
|
||||
message['done'] = True
|
||||
elif not pending:
|
||||
# Only the transaction completing the last call may start the model again.
|
||||
output.append(tool_continuation())
|
||||
else:
|
||||
next_approval = next((item for item in pending if not item.get('approved')), None)
|
||||
if next_approval is not None:
|
||||
next_approval['status'] = 'pending'
|
||||
return output
|
||||
|
||||
|
||||
async def resolve_tool_call_output(
|
||||
chat_id: str,
|
||||
message_id: str,
|
||||
|
|
@ -33,81 +78,98 @@ async def resolve_tool_call_output(
|
|||
detail=ERROR_MESSAGES.ACCESS_PROHIBITED,
|
||||
)
|
||||
|
||||
message = await Chats.get_message_by_id_and_message_id(chat_id, message_id)
|
||||
if not message:
|
||||
raise HTTPException(status_code=status.HTTP_404_NOT_FOUND, detail=ERROR_MESSAGES.NOT_FOUND)
|
||||
async with Chats.edit_message_output(chat_id, message_id) as message:
|
||||
if not message:
|
||||
raise HTTPException(status_code=status.HTTP_404_NOT_FOUND, detail=ERROR_MESSAGES.NOT_FOUND)
|
||||
|
||||
if (message.get('user_id') or chat.user_id) != user.id:
|
||||
raise HTTPException(status_code=403, detail=ERROR_MESSAGES.ACCESS_PROHIBITED)
|
||||
if (message.get('user_id') or chat.user_id) != user.id:
|
||||
raise HTTPException(status_code=403, detail=ERROR_MESSAGES.ACCESS_PROHIBITED)
|
||||
|
||||
output = message.get('output') or []
|
||||
if not isinstance(output, list):
|
||||
raise HTTPException(status_code=status.HTTP_400_BAD_REQUEST, detail='Message has no resolvable output.')
|
||||
output = message.get('output') or []
|
||||
if not isinstance(output, list):
|
||||
raise HTTPException(status_code=status.HTTP_400_BAD_REQUEST, detail='Message has no resolvable output.')
|
||||
|
||||
function_call = next(
|
||||
(
|
||||
item
|
||||
for item in output
|
||||
if item.get('type') == 'function_call' and (item.get('call_id') or item.get('id')) == form_data.call_id
|
||||
),
|
||||
None,
|
||||
)
|
||||
if not function_call:
|
||||
raise HTTPException(status_code=status.HTTP_404_NOT_FOUND, detail='Tool call not found.')
|
||||
function_call.setdefault('call_id', form_data.call_id)
|
||||
tool_name = function_call.get('name')
|
||||
|
||||
if any(
|
||||
item.get('type') == 'function_call_output' and item.get('call_id') == form_data.call_id for item in output
|
||||
) or function_call.get('status') not in {'pending', 'queued', 'requires_approval'}:
|
||||
raise HTTPException(status_code=status.HTTP_409_CONFLICT, detail='Tool call has already been resolved.')
|
||||
|
||||
if form_data.action == 'approve':
|
||||
if tool_name == 'ask_user':
|
||||
raise HTTPException(status_code=status.HTTP_400_BAD_REQUEST, detail='ask_user requires an answer or deny.')
|
||||
function_call['status'] = 'queued'
|
||||
function_call['approved'] = True
|
||||
elif form_data.action == 'reject':
|
||||
function_call['status'] = 'rejected'
|
||||
output.append(
|
||||
{
|
||||
'type': 'function_call_output',
|
||||
'id': f'fco_{form_data.call_id}',
|
||||
'call_id': form_data.call_id,
|
||||
'output': [{'type': 'input_text', 'text': 'Error: tool call rejected by user.'}],
|
||||
'status': 'rejected',
|
||||
}
|
||||
function_call = next(
|
||||
(
|
||||
item
|
||||
for item in output
|
||||
if item.get('type') == 'function_call' and (item.get('call_id') or item.get('id')) == form_data.call_id
|
||||
),
|
||||
None,
|
||||
)
|
||||
else:
|
||||
if tool_name != 'ask_user':
|
||||
raise HTTPException(status_code=status.HTTP_400_BAD_REQUEST, detail='Tool call does not accept answers.')
|
||||
if form_data.answers is None and not form_data.timed_out:
|
||||
raise HTTPException(status_code=status.HTTP_400_BAD_REQUEST, detail='Answers are required for ask_user.')
|
||||
function_call['status'] = 'completed'
|
||||
answer_payload = (
|
||||
{'status': 'cancelled', 'answers': {}, 'timed_out': True}
|
||||
if form_data.timed_out
|
||||
else {'status': 'answered', 'answers': form_data.answers or {}}
|
||||
)
|
||||
output.append(
|
||||
{
|
||||
'type': 'function_call_output',
|
||||
'id': f'fco_{form_data.call_id}',
|
||||
'call_id': form_data.call_id,
|
||||
'output': [{'type': 'input_text', 'text': JSONCodec.dumps(answer_payload)}],
|
||||
'status': 'completed',
|
||||
}
|
||||
if not function_call:
|
||||
raise HTTPException(status_code=status.HTTP_404_NOT_FOUND, detail='Tool call not found.')
|
||||
function_call.setdefault('call_id', form_data.call_id)
|
||||
tool_name = function_call.get('name')
|
||||
|
||||
if (
|
||||
any(
|
||||
item.get('type') == 'function_call_output' and item.get('call_id') == form_data.call_id
|
||||
for item in output
|
||||
)
|
||||
or function_call.get('status') not in {'pending', 'queued', 'requires_approval'}
|
||||
or function_call.get('approved') is True
|
||||
):
|
||||
raise HTTPException(status_code=status.HTTP_409_CONFLICT, detail='Tool call has already been resolved.')
|
||||
|
||||
active = any(
|
||||
item.get('approved') is True or item.get('status') == 'in_progress' for item in pending_tool_calls(output)
|
||||
)
|
||||
|
||||
await Chats.upsert_message_to_chat_by_id_and_message_id(
|
||||
chat_id,
|
||||
message_id,
|
||||
{
|
||||
'done': False,
|
||||
'output': output,
|
||||
},
|
||||
touch=False,
|
||||
)
|
||||
if form_data.action == 'approve':
|
||||
if tool_name == 'ask_user':
|
||||
raise HTTPException(
|
||||
status_code=status.HTTP_400_BAD_REQUEST, detail='ask_user requires an answer or deny.'
|
||||
)
|
||||
function_call['status'] = 'queued'
|
||||
function_call['approved'] = True
|
||||
elif form_data.action == 'reject':
|
||||
function_call['status'] = 'rejected'
|
||||
output.append(
|
||||
{
|
||||
'type': 'function_call_output',
|
||||
'id': f'fco_{form_data.call_id}',
|
||||
'call_id': form_data.call_id,
|
||||
'output': [{'type': 'input_text', 'text': 'Error: tool call rejected by user.'}],
|
||||
'status': 'rejected',
|
||||
}
|
||||
)
|
||||
else:
|
||||
if tool_name != 'ask_user':
|
||||
raise HTTPException(
|
||||
status_code=status.HTTP_400_BAD_REQUEST, detail='Tool call does not accept answers.'
|
||||
)
|
||||
if form_data.answers is None and not form_data.timed_out:
|
||||
raise HTTPException(
|
||||
status_code=status.HTTP_400_BAD_REQUEST, detail='Answers are required for ask_user.'
|
||||
)
|
||||
function_call['status'] = 'completed'
|
||||
answer_payload = (
|
||||
{'status': 'cancelled', 'answers': {}, 'timed_out': True}
|
||||
if form_data.timed_out
|
||||
else {'status': 'answered', 'answers': form_data.answers or {}}
|
||||
)
|
||||
output.append(
|
||||
{
|
||||
'type': 'function_call_output',
|
||||
'id': f'fco_{form_data.call_id}',
|
||||
'call_id': form_data.call_id,
|
||||
'output': [{'type': 'input_text', 'text': JSONCodec.dumps(answer_payload)}],
|
||||
'status': 'completed',
|
||||
}
|
||||
)
|
||||
|
||||
message['done'] = False
|
||||
pending = pending_tool_calls(output)
|
||||
resume = None
|
||||
if not active:
|
||||
if form_data.action == 'approve':
|
||||
resume = 'tools'
|
||||
elif not pending:
|
||||
output.append(tool_continuation())
|
||||
resume = 'model'
|
||||
else:
|
||||
pending[0]['status'] = 'pending'
|
||||
|
||||
event_emitter = await get_event_emitter(
|
||||
{
|
||||
|
|
@ -120,17 +182,7 @@ async def resolve_tool_call_output(
|
|||
if event_emitter:
|
||||
await event_emitter({'type': 'chat:completion', 'data': {'output': output}})
|
||||
|
||||
result_call_ids = {
|
||||
item.get('call_id') for item in output if item.get('type') == 'function_call_output' and item.get('call_id')
|
||||
}
|
||||
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
|
||||
)
|
||||
return {'chat': chat, 'message': message, 'output': output, 'paused': paused}
|
||||
return {'chat': chat, 'output': output, 'resume': resume}
|
||||
|
||||
|
||||
async def build_tool_approval_resume_payload(chat_id: str, message_id: str, chat=None) -> dict:
|
||||
|
|
|
|||
Loading…
Add table
Reference in a new issue