Coverage for open_webui/utils/ask_user.py: 5%

72 statements  

« prev     ^ index     » next       coverage.py v7.15.2, created at 2026-10-07 05:07 +0000

1from collections.abc import Callable 

2 

3from open_webui.utils.json_codec import JSONCodec 

4 

5 

6ASK_USER_NAME = 'ask_user' 

7 

8 

9def get_ask_user_tool_calls(tool_calls: list[dict]) -> tuple[list[dict], str | None]: 

10 ask_user_calls = [ 

11 tool_call for tool_call in tool_calls if tool_call.get('function', {}).get('name') == ASK_USER_NAME 

12 ] 

13 if not ask_user_calls: 

14 return [], None 

15 if len(tool_calls) != 1: 

16 return ( 

17 ask_user_calls, 

18 'Error: ask_user must be the only tool call, so it did not run. Call ask_user on its own.', 

19 ) 

20 if len(ask_user_calls) != 1: 

21 return ask_user_calls, 'Error: only one ask_user call is allowed per turn.' 

22 return ask_user_calls, None 

23 

24 

25def normalize_ask_user_request(arguments: dict) -> dict: 

26 questions = arguments.get('questions') 

27 if not isinstance(questions, list) or not 1 <= len(questions) <= 3: 

28 raise ValueError('ask_user requires 1-3 questions.') 

29 

30 normalized_questions = [] 

31 seen_ids = set() 

32 allow_other = bool(arguments.get('allow_other', True)) 

33 for index, question in enumerate(questions): 

34 if not isinstance(question, dict): 

35 raise ValueError('Each question must be an object.') 

36 

37 question_id = str(question.get('id') or '').strip()[:64] 

38 if not question_id: 

39 raise ValueError('Each question requires a non-empty id.') 

40 if question_id in seen_ids: 

41 raise ValueError(f'Duplicate question id: {question_id}') 

42 seen_ids.add(question_id) 

43 

44 options = question.get('options') 

45 if not isinstance(options, list) or not 2 <= len(options) <= 3: 

46 raise ValueError('Each question requires 2-3 options.') 

47 

48 normalized_options = [] 

49 for option in options: 

50 if not isinstance(option, dict): 

51 raise ValueError('Each option must be an object.') 

52 label = str(option.get('label') or '').strip()[:80] 

53 description = str(option.get('description') or '').strip()[:240] 

54 if not label or not description: 

55 raise ValueError('Each option requires a label and description.') 

56 normalized_options.append({'label': label, 'description': description}) 

57 

58 question_text = str(question.get('question') or '').strip()[:500] 

59 if not question_text: 

60 raise ValueError('Each question requires question text.') 

61 

62 normalized_questions.append( 

63 { 

64 'id': question_id, 

65 'header': str(question.get('header') or '').strip()[:48] or f'Question {index + 1}', 

66 'question': question_text, 

67 'options': normalized_options, 

68 'allow_other': bool(question.get('allow_other', allow_other)), 

69 } 

70 ) 

71 

72 timeout_ms = arguments.get('timeout_ms', 120_000) 

73 if isinstance(timeout_ms, bool) or not isinstance(timeout_ms, int) or not 60_000 <= timeout_ms <= 240_000: 

74 timeout_ms = 120_000 

75 

76 return { 

77 'questions': normalized_questions, 

78 'allow_other': allow_other, 

79 'timeout_ms': timeout_ms, 

80 } 

81 

82 

83def stage_ask_user_tool_calls( 

84 tool_calls: list[dict], 

85 output: list[dict], 

86 make_output_id: Callable[[str], str], 

87) -> tuple[bool, str | None]: 

88 ask_user_calls, error = get_ask_user_tool_calls(tool_calls) 

89 if not ask_user_calls: 

90 return False, None 

91 

92 for tool_call in ask_user_calls: 

93 call_id = tool_call.get('id') or make_output_id('fc') 

94 raw_arguments = tool_call.get('function', {}).get('arguments', '{}') 

95 arguments = raw_arguments 

96 

97 if not error: 

98 try: 

99 parsed_arguments = JSONCodec.loads(raw_arguments or '{}') 

100 if not isinstance(parsed_arguments, dict): 

101 raise ValueError('ask_user arguments must be an object.') 

102 arguments = JSONCodec.dumps(normalize_ask_user_request(parsed_arguments)) 

103 except (JSONCodec.JSONDecodeError, TypeError, ValueError) as exc: 

104 error = f'Error: {exc}' 

105 

106 item = { 

107 'type': 'function_call', 

108 'id': call_id or make_output_id('fc'), 

109 'call_id': call_id, 

110 'name': ASK_USER_NAME, 

111 'arguments': arguments, 

112 'status': 'completed' if error else 'pending', 

113 } 

114 

115 existing_item = next( 

116 ( 

117 existing 

118 for existing in output 

119 if existing.get('type') == 'function_call' 

120 and ( 

121 existing.get('call_id') == call_id 

122 or existing.get('id') == tool_call.get('id') 

123 or ( 

124 not existing.get('call_id') 

125 and existing.get('name') == ASK_USER_NAME 

126 and existing.get('status') not in {'rejected', 'failed'} 

127 ) 

128 ) 

129 ), 

130 None, 

131 ) 

132 if existing_item: 

133 existing_item.update(item) 

134 else: 

135 output.append(item) 

136 

137 # Every invalid call needs its own result, or the UI waits on it forever. 

138 if error: 

139 output.append( 

140 { 

141 'type': 'function_call_output', 

142 'id': make_output_id('fco'), 

143 'call_id': call_id, 

144 'output': [{'type': 'input_text', 'text': error}], 

145 'status': 'completed', 

146 } 

147 ) 

148 

149 return True, error