Coverage for open_webui/utils/headers.py: 60%
73 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
1import logging
2import time
3from string import punctuation
4from typing import Any, Optional
5from urllib.parse import quote
7import jwt
8from open_webui.env import (
9 FORWARD_USER_INFO_HEADER_AUTH_TYPE,
10 FORWARD_USER_INFO_HEADER_JWT,
11 FORWARD_USER_INFO_HEADER_JWT_EXPIRES_SECONDS,
12 FORWARD_USER_INFO_HEADER_JWT_SECRET,
13 FORWARD_USER_INFO_HEADER_USER_EMAIL,
14 FORWARD_USER_INFO_HEADER_USER_ID,
15 FORWARD_USER_INFO_HEADER_USER_NAME,
16 FORWARD_USER_INFO_HEADER_USER_ROLE,
17)
18from open_webui.models.groups import Groups
20log = logging.getLogger(__name__)
22USER_GROUPS_PLACEHOLDERS = ('{{USER_GROUPS}}', '{{USER_GROUP_IDS}}')
25def normalize_bearer_token(token: Any) -> str:
26 return token.strip() if isinstance(token, str) else token or ''
29def bearer_auth_header(token: Any) -> dict[str, str]:
30 token = normalize_bearer_token(token)
31 return {'Authorization': f'Bearer {token}'} if token else {}
34def get_json_bearer_headers(token: Any = '') -> dict[str, str]:
35 return {'Content-Type': 'application/json', **bearer_auth_header(token)}
38def _mint_forward_user_jwt(user: Any) -> str:
39 now = int(time.time())
40 payload = {
41 'sub': str(user.id),
42 'email': str(user.email),
43 'name': str(user.name),
44 'role': str(user.role),
45 'iss': 'open-webui',
46 'iat': now,
47 'exp': now + FORWARD_USER_INFO_HEADER_JWT_EXPIRES_SECONDS,
48 }
49 return jwt.encode(payload, FORWARD_USER_INFO_HEADER_JWT_SECRET, algorithm='HS256')
52def include_user_info_headers(headers: dict, user: Optional[Any] = None, *, request=None) -> dict:
53 """
54 Forward user identity to external backends: signed JWT in
55 FORWARD_USER_INFO_HEADER_JWT if FORWARD_USER_INFO_HEADER_JWT_SECRET is set;
56 otherwise the legacy X-OpenWebUI-User-* headers.
57 Include the verified incoming auth type when a request provides it.
58 """
59 if user is None:
60 return headers
62 auth_type = getattr(getattr(request, 'state', None), 'auth_type', None)
63 if auth_type in ('api_key', 'jwt'):
64 headers = {**headers, FORWARD_USER_INFO_HEADER_AUTH_TYPE: auth_type}
66 if FORWARD_USER_INFO_HEADER_JWT_SECRET:
67 try:
68 token = _mint_forward_user_jwt(user)
69 return {**headers, FORWARD_USER_INFO_HEADER_JWT: token}
70 except Exception:
71 log.exception(
72 'Failed to mint %s; falling back to plain user-info headers.',
73 FORWARD_USER_INFO_HEADER_JWT,
74 )
76 return {
77 **headers,
78 FORWARD_USER_INFO_HEADER_USER_NAME: quote(user.name.strip(), safe=' '),
79 FORWARD_USER_INFO_HEADER_USER_ID: user.id,
80 FORWARD_USER_INFO_HEADER_USER_EMAIL: user.email.strip(),
81 FORWARD_USER_INFO_HEADER_USER_ROLE: user.role,
82 }
85def custom_headers_require_user_groups(custom_headers: Optional[dict]) -> bool:
86 if not custom_headers or not isinstance(custom_headers, dict): 86 ↛ 87line 86 didn't jump to line 87 because the condition on line 86 was never true
87 return False
88 return any(
89 placeholder in str(value) for value in custom_headers.values() for placeholder in USER_GROUPS_PLACEHOLDERS
90 )
93async def get_user_groups_for_custom_headers(
94 custom_headers: Optional[dict], user: Optional[Any] = None
95) -> Optional[list]:
96 """Fetch the user's groups only when a header value actually references a groups placeholder."""
97 if user is None or not custom_headers_require_user_groups(custom_headers): 97 ↛ 100line 97 didn't jump to line 100 because the condition on line 97 was always true
98 return None
100 try:
101 return await Groups.get_groups_by_member_id(user.id)
102 except Exception:
103 log.exception('Failed to resolve user groups for custom headers')
104 return None
107async def get_custom_headers(custom_headers: dict, user=None, metadata: dict = None, request=None) -> dict:
108 user_groups = await get_user_groups_for_custom_headers(custom_headers, user)
109 return parse_custom_headers(custom_headers, user, metadata, request=request, user_groups=user_groups)
112def parse_custom_headers(
113 custom_headers: dict, user=None, metadata: dict = None, request=None, user_groups: Optional[list] = None
114) -> dict:
115 if not custom_headers or not isinstance(custom_headers, dict): 115 ↛ 116line 115 didn't jump to line 116 because the condition on line 115 was never true
116 return {}
118 metadata = metadata or {}
120 # UA from the live request; fall back to metadata for detached RAG/tool calls.
121 user_agent = ''
122 if request is not None: 122 ↛ 123line 122 didn't jump to line 123 because the condition on line 122 was never true
123 try:
124 user_agent = request.headers.get('user-agent', '') or ''
125 except Exception:
126 user_agent = ''
127 if not user_agent: 127 ↛ 131line 127 didn't jump to line 131 because the condition on line 127 was always true
128 user_agent = metadata.get('user_agent', '') or ''
130 # Extract user_message info for tree mapping
131 user_message = metadata.get('user_message') or {}
132 user_message_id = metadata.get('user_message_id', '') or (user_message.get('id', '') if user_message else '')
133 user_message_parent_id = user_message.get('parentId', '') if user_message else ''
135 template_vars = {
136 '{{CHAT_ID}}': metadata.get('chat_id', '') or '',
137 '{{MESSAGE_ID}}': metadata.get('message_id', '') or '',
138 '{{USER_MESSAGE_ID}}': user_message_id or '',
139 '{{USER_MESSAGE_PARENT_ID}}': user_message_parent_id or '',
140 '{{FILE_ID}}': metadata.get('file_id', '') or '',
141 '{{FILE_NAME}}': metadata.get('file_name', '') or '',
142 '{{FILE_CONTENT_TYPE}}': metadata.get('file_content_type', '') or '',
143 '{{TASK}}': metadata.get('task', '') or '',
144 '{{USER_ID}}': (user.id if user else '') or '',
145 '{{USER_NAME}}': (user.name.strip() if user else '') or '',
146 '{{USER_EMAIL}}': (user.email.strip() if user else '') or '',
147 '{{USER_ROLE}}': (user.role if user else '') or '',
148 '{{USER_GROUPS}}': ','.join(group.name.strip() for group in user_groups) if user_groups else '',
149 '{{USER_GROUP_IDS}}': ','.join(group.id for group in user_groups) if user_groups else '',
150 '{{USER_AGENT}}': user_agent,
151 '{{AUTH_TYPE}}': getattr(getattr(request, 'state', None), 'auth_type', None) or '',
152 }
154 parsed_headers = {}
155 for key, value in custom_headers.items():
156 if not isinstance(value, str):
157 value = str(value)
158 for token, val in template_vars.items():
159 value = value.replace(token, val)
160 # Encode Unicode and controls after substitution; preserve ASCII header syntax and existing escapes.
161 parsed_headers[key] = quote(value, safe=punctuation + ' \t')
163 return parsed_headers