This commit is contained in:
Timothy Jaeryang Baek
2026-08-13 21:36:41 -06:00
parent b4738d1a2e
commit fa94a5ab24
3 changed files with 132 additions and 67 deletions
+42 -8
View File
@@ -1851,14 +1851,6 @@ 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)
if form_data.action != 'approve' and resolution['paused']:
return {
'status': True,
'chat_id': id,
'message_id': message_id,
'paused': True,
}
payload = await build_tool_approval_resume_payload(id, message_id, chat=resolution['chat'])
result = await chat_completion(request, payload, user)
return {
@@ -2119,6 +2111,7 @@ async def list_tasks_by_chat_id_endpoint(request: Request, chat_id: str, user=De
@app.post('/api/tasks/chat/{chat_id:path}/stop')
async def stop_tasks_by_chat_id_endpoint(request: Request, chat_id: str, user=Depends(get_verified_user)):
socket_id = get_temporary_chat_session_id(chat_id)
chat = None
if socket_id:
owner_id = get_user_id_from_session_pool(socket_id)
if owner_id != user.id and user.role != 'admin':
@@ -2128,6 +2121,47 @@ async def stop_tasks_by_chat_id_endpoint(request: Request, chat_id: str, user=De
if chat is None or (chat.user_id != user.id and user.role != 'admin'):
raise HTTPException(status_code=status.HTTP_404_NOT_FOUND, detail=ERROR_MESSAGES.NOT_FOUND)
result = await stop_item_tasks(request.app.state.redis, chat_id)
if not socket_id and str(result.get('message', '')).startswith('No tasks found'):
messages_map = await Chats.get_messages_map_by_chat_id(chat_id) or {}
for message_id, message in messages_map.items():
if message.get('role') != 'assistant' or message.get('done') is not False:
continue
output = message.get('output')
if isinstance(output, list):
for item in output:
if item.get('type') == 'function_call' and item.get('status') in {
'pending',
'queued',
'requires_approval',
}:
item['status'] = 'rejected'
item.pop('approved', None)
await Chats.upsert_message_to_chat_by_id_and_message_id(
chat_id,
message_id,
{'done': True, **({'output': output} if isinstance(output, list) else {})},
touch=False,
)
result = {
'status': True,
'message': 'Finalized pending approval message.',
}
event_emitter = await get_event_emitter(
{
'user_id': chat.user_id,
'chat_id': chat_id,
'message_id': message_id,
},
update_db=False,
)
if event_emitter:
await event_emitter({'type': 'chat:completion', 'data': {'done': True, 'output': output}})
await event_emitter({'type': 'chat:tasks:cancel'})
return result
+35 -13
View File
@@ -2063,16 +2063,6 @@ def process_messages_with_output(
return processed
def strip_compaction_fields(messages: list[dict]) -> list[dict]:
stripped = []
for message in messages:
clean = dict(message)
for key in ('id', 'files', 'output', 'contextSummary', 'context_summary', 'usage'):
clean.pop(key, None)
stripped.append(clean)
return stripped
def sanitize_tool_pairs(messages: list[dict]) -> list[dict]:
tool_result_ids = {
message.get('tool_call_id')
@@ -2331,8 +2321,6 @@ async def process_chat_payload(request, form_data, user, metadata, model):
except Exception:
log.exception('Context compaction failed; continuing with full chat history')
form_data['messages'] = strip_compaction_fields(form_data.get('messages', []))
# Process messages with OR-aligned output items for clean LLM messages
form_data['messages'] = process_messages_with_output(
form_data.get('messages', []),
@@ -3109,6 +3097,20 @@ async def drain_approved_tool_calls(request, form_data, user, model, metadata) -
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
return False
event_emitter, event_caller = await get_event_emitter_and_caller(metadata)
@@ -3146,6 +3148,21 @@ async def drain_approved_tool_calls(request, form_data, user, model, 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 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')
}
paused = any(
item.get('type') == 'function_call'
and item.get('call_id')
@@ -3207,6 +3224,7 @@ async def pause_for_tool_approval(chat_id: str, message_id: str, output: list[di
result_call_ids = {
item.get('call_id') for item in output if item.get('type') == 'function_call_output' and item.get('call_id')
}
has_pending_approval = False
for item in output:
if item.get('type') == 'function_call' and not item.get('call_id') and item.get('id'):
item['call_id'] = item['id']
@@ -3217,7 +3235,11 @@ async def pause_for_tool_approval(chat_id: str, message_id: str, output: list[di
and item.get('call_id') not in result_call_ids
and item.get('status') != 'rejected'
):
item['status'] = 'pending'
if not has_pending_approval:
item['status'] = 'pending'
has_pending_approval = True
elif item.get('status') == 'in_progress':
item['status'] = 'queued'
await Chats.upsert_message_to_chat_by_id_and_message_id(
chat_id,
+55 -46
View File
@@ -293,6 +293,7 @@ def convert_output_to_messages(
pending_reasoning = [] # Only populated when reasoning_format == 'reasoning_content'
pending_reasoning_details = []
pending_tool_image_urls = []
pending_tool_outputs = []
completed_call_ids = {
item.get('call_id')
for item in output
@@ -347,9 +348,60 @@ def convert_output_to_messages(
)
pending_tool_image_urls = []
def flush_tool_outputs():
nonlocal pending_tool_outputs
if not pending_tool_outputs:
return
flush_pending()
for output_item in pending_tool_outputs:
output_parts = output_item.get('output', [])
content = ''
image_urls = []
for part in output_parts:
if part.get('type') == 'input_text':
output_text = part.get('text', '')
content += str(output_text) if not isinstance(output_text, str) else output_text
elif part.get('type') == 'input_image':
url = part.get('image_url', '')
if url:
image_urls.append(url)
if flatten_tool_images:
messages.append(
{
'role': 'tool',
'tool_call_id': output_item.get('call_id', ''),
'content': content,
}
)
pending_tool_image_urls.extend(image_urls)
elif image_urls:
messages.append(
{
'role': 'tool',
'tool_call_id': output_item.get('call_id', ''),
'content': [
{'type': 'input_text', 'text': content},
*[{'type': 'input_image', 'image_url': url} for url in image_urls],
],
}
)
else:
messages.append(
{
'role': 'tool',
'tool_call_id': output_item.get('call_id', ''),
'content': content,
}
)
pending_tool_outputs = []
for item in output:
item_type = item.get('type', '')
if item_type != 'function_call_output':
if item_type not in {'function_call', 'function_call_output'}:
flush_tool_outputs()
flush_tool_images()
if item_type == 'message':
@@ -386,51 +438,7 @@ def convert_output_to_messages(
if item.get('call_id') not in function_call_ids:
continue
# Flush any pending content/tool_calls before adding tool result
flush_pending()
# Extract text and images from output content parts
output_parts = item.get('output', [])
content = ''
image_urls = []
for part in output_parts:
if part.get('type') == 'input_text':
output_text = part.get('text', '')
content += str(output_text) if not isinstance(output_text, str) else output_text
elif part.get('type') == 'input_image':
url = part.get('image_url', '')
if url:
image_urls.append(url)
if flatten_tool_images:
messages.append(
{
'role': 'tool',
'tool_call_id': item.get('call_id', ''),
'content': content,
}
)
if item.get('call_id') in function_call_ids:
pending_tool_image_urls.extend(image_urls)
elif image_urls:
messages.append(
{
'role': 'tool',
'tool_call_id': item.get('call_id', ''),
'content': [
{'type': 'input_text', 'text': content},
*[{'type': 'input_image', 'image_url': url} for url in image_urls],
],
}
)
else:
messages.append(
{
'role': 'tool',
'tool_call_id': item.get('call_id', ''),
'content': content,
}
)
pending_tool_outputs.append(item)
elif item_type == 'reasoning':
reasoning_details = item.get('reasoning_details') if raw else None
@@ -484,6 +492,7 @@ def convert_output_to_messages(
pass
# Flush remaining content/tool_calls
flush_tool_outputs()
flush_tool_images()
flush_pending()