Coverage for open_webui/utils/tool_approval.py: 28%
79 statements
« prev ^ index » next coverage.py v7.15.2, created at 2026-10-07 05:07 +0000
« prev ^ index » next coverage.py v7.15.2, created at 2026-10-07 05:07 +0000
1from typing import Any, Literal
3from fastapi import HTTPException, status
4from pydantic import BaseModel
5from sqlalchemy.ext.asyncio import AsyncSession
7from open_webui.constants import ERROR_MESSAGES
8from open_webui.models.chats import Chats
9from open_webui.socket.main import get_event_emitter
10from open_webui.utils.json_codec import JSONCodec
13class ResolveToolCallForm(BaseModel):
14 call_id: str
15 action: Literal['approve', 'reject', 'answer']
16 answers: Any | None = None
17 timed_out: bool = False
20async def resolve_tool_call_output(
21 chat_id: str,
22 message_id: str,
23 form_data: ResolveToolCallForm,
24 user,
25 db: AsyncSession | None = None,
26) -> dict:
27 chat = await Chats.get_chat_by_id(chat_id, db=db)
28 if not chat or (chat.user_id != user.id and user.role != 'admin'):
29 raise HTTPException(
30 status_code=status.HTTP_401_UNAUTHORIZED,
31 detail=ERROR_MESSAGES.ACCESS_PROHIBITED,
32 )
34 message = await Chats.get_message_by_id_and_message_id(chat_id, message_id)
35 if not message:
36 raise HTTPException(status_code=status.HTTP_404_NOT_FOUND, detail=ERROR_MESSAGES.NOT_FOUND)
38 output = message.get('output') or []
39 if not isinstance(output, list): 39 ↛ 40line 39 didn't jump to line 40 because the condition on line 39 was never true
40 raise HTTPException(status_code=status.HTTP_400_BAD_REQUEST, detail='Message has no resolvable output.')
42 function_call = next(
43 (
44 item
45 for item in output
46 if item.get('type') == 'function_call' and (item.get('call_id') or item.get('id')) == form_data.call_id
47 ),
48 None,
49 )
50 if not function_call: 50 ↛ 52line 50 didn't jump to line 52 because the condition on line 50 was always true
51 raise HTTPException(status_code=status.HTTP_404_NOT_FOUND, detail='Tool call not found.')
52 function_call.setdefault('call_id', form_data.call_id)
53 tool_name = function_call.get('name')
55 if any(
56 item.get('type') == 'function_call_output' and item.get('call_id') == form_data.call_id for item in output
57 ) or function_call.get('status') not in {'pending', 'queued', 'requires_approval'}:
58 raise HTTPException(status_code=status.HTTP_409_CONFLICT, detail='Tool call has already been resolved.')
60 if form_data.action == 'approve':
61 if tool_name == 'ask_user':
62 raise HTTPException(status_code=status.HTTP_400_BAD_REQUEST, detail='ask_user requires an answer or deny.')
63 function_call['status'] = 'queued'
64 function_call['approved'] = True
65 elif form_data.action == 'reject':
66 function_call['status'] = 'rejected'
67 output.append(
68 {
69 'type': 'function_call_output',
70 'id': f'fco_{form_data.call_id}',
71 'call_id': form_data.call_id,
72 'output': [{'type': 'input_text', 'text': 'Error: tool call rejected by user.'}],
73 'status': 'rejected',
74 }
75 )
76 else:
77 if tool_name != 'ask_user':
78 raise HTTPException(status_code=status.HTTP_400_BAD_REQUEST, detail='Tool call does not accept answers.')
79 if form_data.answers is None and not form_data.timed_out:
80 raise HTTPException(status_code=status.HTTP_400_BAD_REQUEST, detail='Answers are required for ask_user.')
81 function_call['status'] = 'completed'
82 answer_payload = (
83 {'status': 'cancelled', 'answers': {}, 'timed_out': True}
84 if form_data.timed_out
85 else {'status': 'answered', 'answers': form_data.answers or {}}
86 )
87 output.append(
88 {
89 'type': 'function_call_output',
90 'id': f'fco_{form_data.call_id}',
91 'call_id': form_data.call_id,
92 'output': [{'type': 'input_text', 'text': JSONCodec.dumps(answer_payload)}],
93 'status': 'completed',
94 }
95 )
97 await Chats.upsert_message_to_chat_by_id_and_message_id(
98 chat_id,
99 message_id,
100 {
101 'done': False,
102 'output': output,
103 },
104 touch=False,
105 )
107 event_emitter = await get_event_emitter(
108 {
109 'user_id': chat.user_id,
110 'chat_id': chat_id,
111 'message_id': message_id,
112 },
113 update_db=False,
114 )
115 if event_emitter:
116 await event_emitter({'type': 'chat:completion', 'data': {'output': output}})
118 result_call_ids = {
119 item.get('call_id') for item in output if item.get('type') == 'function_call_output' and item.get('call_id')
120 }
121 paused = any(
122 item.get('type') == 'function_call'
123 and item.get('call_id')
124 and item.get('status') in {'pending', 'queued', 'requires_approval'}
125 and item.get('call_id') not in result_call_ids
126 for item in output
127 )
128 return {'chat': chat, 'message': message, 'output': output, 'paused': paused}
131async def build_tool_approval_resume_payload(chat_id: str, message_id: str, chat=None) -> dict:
132 chat = chat or await Chats.get_chat_by_id(chat_id)
133 if not chat:
134 raise HTTPException(status_code=status.HTTP_404_NOT_FOUND, detail=ERROR_MESSAGES.NOT_FOUND)
136 assistant_message = await Chats.get_message_by_id_and_message_id(chat_id, message_id)
137 if not assistant_message:
138 raise HTTPException(status_code=status.HTTP_404_NOT_FOUND, detail=ERROR_MESSAGES.NOT_FOUND)
140 user_message_id = assistant_message.get('parentId')
141 user_message = await Chats.get_message_by_id_and_message_id(chat_id, user_message_id) if user_message_id else None
142 if not user_message:
143 raise HTTPException(status_code=status.HTTP_409_CONFLICT, detail='Tool call parent message is missing.')
145 chat_data = chat.chat or {}
146 message_meta = assistant_message.get('meta') if isinstance(assistant_message.get('meta'), dict) else {}
147 chat_params = chat_data.get('params') if isinstance(chat_data.get('params'), dict) else {}
148 params = {
149 **chat_params,
150 **(message_meta.get('params') if isinstance(message_meta.get('params'), dict) else {}),
151 }
152 current_approval_mode = chat_params.get('tool_approval_mode')
153 if current_approval_mode in {'ask', 'full'}:
154 params['tool_approval_mode'] = current_approval_mode
155 if 'tool_approval_mode' not in params:
156 params['tool_approval_mode'] = 'ask'
158 model_id = assistant_message.get('model') or next(iter(chat_data.get('models') or []), None)
159 if not model_id:
160 raise HTTPException(status_code=status.HTTP_409_CONFLICT, detail='Tool call message model is missing.')
162 messages = []
163 if params.get('system'):
164 messages.append({'role': 'system', 'content': params.get('system')})
166 return {
167 'stream': params.get('stream_response', True),
168 'model': model_id,
169 'messages': messages,
170 'params': params,
171 'files': message_meta.get('files') or chat_data.get('files') or None,
172 'filter_ids': message_meta.get('filter_ids') or None,
173 'tool_ids': message_meta.get('tool_ids') or None,
174 'skill_ids': message_meta.get('skill_ids') or None,
175 'terminal_id': message_meta.get('terminal_id') or None,
176 'tool_servers': message_meta.get('tool_servers') or None,
177 'features': message_meta.get('features') or {},
178 'variables': message_meta.get('variables') or {},
179 'chat_variables': chat.variables,
180 'session_id': message_meta.get('session_id'),
181 'chat_id': chat_id,
182 'id': message_id,
183 'parent_id': user_message.get('parentId'),
184 'user_message': user_message,
185 'assistant_message_id': message_id,
186 }