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

1from typing import Any, Literal 

2 

3from fastapi import HTTPException, status 

4from pydantic import BaseModel 

5from sqlalchemy.ext.asyncio import AsyncSession 

6 

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 

11 

12 

13class ResolveToolCallForm(BaseModel): 

14 call_id: str 

15 action: Literal['approve', 'reject', 'answer'] 

16 answers: Any | None = None 

17 timed_out: bool = False 

18 

19 

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 ) 

33 

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) 

37 

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.') 

41 

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') 

54 

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.') 

59 

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 ) 

96 

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 ) 

106 

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}}) 

117 

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} 

129 

130 

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) 

135 

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) 

139 

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.') 

144 

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' 

157 

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.') 

161 

162 messages = [] 

163 if params.get('system'): 

164 messages.append({'role': 'system', 'content': params.get('system')}) 

165 

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 }