Coverage for open_webui/utils/chat_variables.py: 16%
195 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 __future__ import annotations
3import re
4from typing import Any
6from open_webui.utils.json_codec import JSONCodec
8CHAT_VARIABLE_KEY_RE = re.compile(r'^[a-z][a-z0-9_]*$')
9CHAT_VARIABLE_ANY_RE = re.compile(r'{{\s*chat\.variables\.([^\s|}]+)(?:\s*\|\s*([^}]*))?\s*}}')
10USER_VARIABLE_ANY_RE = re.compile(r'{{\s*user\.variables\.([^\s|}]+)(?:\s*\|\s*([^}]*))?\s*}}')
11MAX_VARIABLE_VALUE_LENGTH = 20_000
12MAX_VARIABLES_JSON_LENGTH = 100_000
15class ChatVariablesError(ValueError):
16 pass
19def split_properties(value: str, delimiter: str) -> list[str]:
20 result: list[str] = []
21 current = ''
22 depth = 0
23 in_string = False
24 escape_next = False
26 for char in value:
27 if escape_next:
28 current += char
29 escape_next = False
30 continue
32 if char == '\\':
33 current += char
34 escape_next = True
35 continue
37 if char == '"' and not escape_next:
38 in_string = not in_string
39 current += char
40 continue
42 if not in_string:
43 if char in ('{', '['):
44 depth += 1
45 elif char in ('}', ']'):
46 depth -= 1
48 if char == delimiter and depth == 0:
49 result.append(current.strip())
50 current = ''
51 continue
53 current += char
55 if current.strip():
56 result.append(current.strip())
58 return result
61def parse_json_value(value: str) -> Any:
62 if value.startswith('"') and value.endswith('"'):
63 return value[1:-1]
65 if re.match(r'^[\[{]', value):
66 try:
67 return JSONCodec.loads(value)
68 except JSONCodec.JSONDecodeError:
69 return value
71 return value
74def parse_variable_definition(definition: str) -> dict[str, Any]:
75 parts = split_properties(definition, ':')
76 if not parts:
77 return {'type': 'text'}
79 first_part, *property_parts = parts
80 field_type = first_part[5:] if first_part.startswith('type=') else first_part
81 field_type = field_type.strip() or 'text'
82 properties: dict[str, Any] = {}
84 for part in property_parts:
85 trimmed = part.strip()
86 if not trimmed:
87 continue
89 equals_parts = split_properties(trimmed, '=')
90 if len(equals_parts) == 1:
91 properties[equals_parts[0].strip()] = True
92 continue
94 property_name, *value_parts = equals_parts
95 properties[property_name.strip()] = parse_json_value('='.join(value_parts).strip())
97 return {'type': field_type, **properties}
100def _safe_field(key: str, definition: dict[str, Any]) -> dict[str, Any]:
101 allowed_keys = {
102 'default',
103 'label',
104 'max',
105 'maxlength',
106 'min',
107 'minlength',
108 'options',
109 'placeholder',
110 'required',
111 'step',
112 'type',
113 }
114 field = {'key': key}
115 for field_key in sorted(allowed_keys):
116 if field_key in definition:
117 field[field_key] = definition[field_key]
119 field.setdefault('type', 'text')
120 if field.get('type') == 'select' and not isinstance(field.get('options'), list):
121 field['options'] = []
122 field['required'] = bool(field.get('required', False))
124 return field
127def get_chat_variables_schema(system_prompt: str | None) -> dict[str, list[dict[str, Any]]] | None:
128 if not system_prompt: 128 ↛ 131line 128 didn't jump to line 131 because the condition on line 128 was always true
129 return None
131 try:
132 fields_by_key = collect_chat_variable_fields(system_prompt)
133 except ChatVariablesError:
134 fields_by_key = {}
136 if not fields_by_key:
137 return None
139 return {'fields': list(fields_by_key.values())}
142def collect_chat_variable_fields(system_prompt: str | None) -> dict[str, dict[str, Any]]:
143 fields_by_key: dict[str, dict[str, Any]] = {}
144 if not system_prompt:
145 return fields_by_key
147 typed_fields_by_key: dict[str, dict[str, Any]] = {}
148 for match in CHAT_VARIABLE_ANY_RE.finditer(system_prompt):
149 key = match.group(1).strip()
150 definition = match.group(2)
151 if not CHAT_VARIABLE_KEY_RE.match(key):
152 raise ChatVariablesError(f'Invalid chat variable key: {key}')
154 if definition is None or not definition.strip():
155 fields_by_key.setdefault(key, _safe_field(key, {'type': 'text'}))
156 continue
158 field = _safe_field(key, parse_variable_definition(definition.strip()))
160 if field.get('type') == 'select' and not field.get('options'):
161 raise ChatVariablesError(f'Chat variable {key} select needs options.')
162 previous = typed_fields_by_key.get(key)
163 if previous and previous != field:
164 raise ChatVariablesError(f'Chat variable {key} has conflicting definitions.')
165 typed_fields_by_key[key] = field
166 fields_by_key[key] = field
168 return fields_by_key
171def normalize_chat_variables(variables: Any) -> dict[str, Any]:
172 if not isinstance(variables, dict):
173 return {}
174 return variables
177def normalize_user_variables(variables: Any) -> dict[str, str]:
178 if not isinstance(variables, dict): 178 ↛ 179line 178 didn't jump to line 179 because the condition on line 178 was never true
179 return {}
180 return {key: value for key, value in variables.items() if isinstance(key, str) and isinstance(value, str)}
183def validate_user_variables(variables: Any) -> dict[str, str]:
184 if not isinstance(variables, dict): 184 ↛ 185line 184 didn't jump to line 185 because the condition on line 184 was never true
185 raise ChatVariablesError('User variables must be an object.')
187 try:
188 if len(JSONCodec.dumps(variables)) > MAX_VARIABLES_JSON_LENGTH: 188 ↛ 189line 188 didn't jump to line 189 because the condition on line 188 was never true
189 raise ChatVariablesError('User variables are too large.')
190 except TypeError:
191 raise ChatVariablesError('User variables must be JSON serializable.')
193 validated: dict[str, str] = {}
194 for key, value in variables.items():
195 if not isinstance(key, str) or not CHAT_VARIABLE_KEY_RE.match(key):
196 raise ChatVariablesError(f'Invalid user variable key: {key}')
197 if not isinstance(value, str): 197 ↛ 199line 197 didn't jump to line 199 because the condition on line 197 was always true
198 raise ChatVariablesError(f'User variable must be a string: {key}')
199 value = value.replace('\r\n', '\n')
200 if len(value) > MAX_VARIABLE_VALUE_LENGTH:
201 raise ChatVariablesError(f'User variable is too long: {key}')
202 validated[key] = value
204 return validated
207def validate_chat_variables(
208 system_prompt: str | None,
209 variables: Any,
210 *,
211 required: bool = True,
212) -> dict[str, Any]:
213 field_map = collect_chat_variable_fields(system_prompt)
214 variables = normalize_chat_variables(variables)
216 try:
217 if len(JSONCodec.dumps(variables)) > MAX_VARIABLES_JSON_LENGTH:
218 raise ChatVariablesError('Chat variables are too large.')
219 except TypeError:
220 raise ChatVariablesError('Chat variables must be JSON serializable.')
222 validated: dict[str, Any] = {}
223 for key, field in field_map.items():
224 has_value = key in variables and variables[key] not in (None, '')
225 value = variables.get(key)
227 if not has_value:
228 if field.get('default') not in (None, ''):
229 value = field.get('default')
230 has_value = True
231 elif required and field.get('required'):
232 label = field.get('label') or key
233 raise ChatVariablesError(f'Missing required chat variable: {label}')
234 else:
235 value = ''
237 if field.get('type') == 'select':
238 options = field.get('options') or []
239 if has_value and value not in options:
240 label = field.get('label') or key
241 raise ChatVariablesError(f'Invalid value for chat variable: {label}')
243 if isinstance(value, str):
244 value = value.replace('\r\n', '\n')
245 if len(value) > MAX_VARIABLE_VALUE_LENGTH:
246 label = field.get('label') or key
247 raise ChatVariablesError(f'Chat variable is too long: {label}')
249 validated[key] = value
251 return validated
254def render_chat_variables(
255 system_prompt: str | None,
256 variables: Any,
257 *,
258 required: bool = True,
259) -> str | None:
260 if not system_prompt:
261 return system_prompt
263 try:
264 validated = validate_chat_variables(system_prompt, variables, required=required)
265 except ChatVariablesError:
266 validated = {}
268 def replace(match: re.Match) -> str:
269 key = match.group(1).strip()
270 value = validated.get(key, '')
271 return '' if value is None else str(value)
273 return CHAT_VARIABLE_ANY_RE.sub(replace, system_prompt)
276def render_user_variables(system_prompt: str | None, variables: Any) -> str | None:
277 if not system_prompt:
278 return system_prompt
280 variables = normalize_user_variables(variables)
282 def replace(match: re.Match) -> str:
283 key = match.group(1).strip()
284 if not CHAT_VARIABLE_KEY_RE.match(key):
285 return ''
286 return variables.get(key, '')
288 return USER_VARIABLE_ANY_RE.sub(replace, system_prompt)