mirror of
https://github.com/open-webui/open-webui.git
synced 2026-10-06 02:48:04 +00:00
refac
This commit is contained in:
parent
19957bc19b
commit
bc2416c5db
4 changed files with 137 additions and 121 deletions
|
|
@ -22,8 +22,6 @@ from open_webui.env import (
|
|||
AIOHTTP_CLIENT_TIMEOUT_MODEL_LIST,
|
||||
AIOHTTP_FILE_STREAM_CHUNK_SIZE,
|
||||
BYPASS_MODEL_ACCESS_CONTROL,
|
||||
ENABLE_FORWARD_USER_INFO_HEADERS,
|
||||
FORWARD_SESSION_INFO_HEADER_CHAT_ID,
|
||||
MODELS_CACHE_TTL,
|
||||
REDIS_KEY_PREFIX,
|
||||
)
|
||||
|
|
@ -36,7 +34,7 @@ from open_webui.models.models import Models
|
|||
from open_webui.models.users import UserModel
|
||||
from open_webui.utils.access_control import check_model_access
|
||||
from open_webui.utils.auth import get_admin_user, get_verified_user
|
||||
from open_webui.utils.headers import get_custom_headers, include_user_info_headers
|
||||
from open_webui.utils.headers import get_headers_and_cookies
|
||||
from open_webui.utils.json_codec import JSONCodec
|
||||
from open_webui.utils.misc import calculate_sha256
|
||||
from open_webui.utils.model_ids import strip_provider_model_prefix
|
||||
|
|
@ -67,24 +65,21 @@ def _clean_proxy_headers(raw_headers) -> dict:
|
|||
|
||||
|
||||
async def send_get_request(
|
||||
url: str,
|
||||
key: str | None = None,
|
||||
user: UserModel | None = None,
|
||||
request: Request = None,
|
||||
url=None,
|
||||
key=None,
|
||||
user: UserModel = None,
|
||||
config=None,
|
||||
):
|
||||
"""Issue a GET request to an Ollama backend and return JSON, or *None* on failure."""
|
||||
try:
|
||||
session = await get_session()
|
||||
headers: dict = {
|
||||
'Content-Type': 'application/json',
|
||||
}
|
||||
if key:
|
||||
headers['Authorization'] = f'Bearer {key}'
|
||||
if ENABLE_FORWARD_USER_INFO_HEADERS and user:
|
||||
headers = include_user_info_headers(headers, user)
|
||||
headers, cookies = await get_headers_and_cookies(request, url, key, config, user=user)
|
||||
|
||||
async with session.get(
|
||||
url,
|
||||
headers=headers,
|
||||
cookies=cookies,
|
||||
ssl=AIOHTTP_CLIENT_SESSION_SSL,
|
||||
timeout=_MODEL_LIST_TIMEOUT,
|
||||
) as r:
|
||||
|
|
@ -114,25 +109,14 @@ async def send_request(
|
|||
try:
|
||||
session = await get_session()
|
||||
|
||||
headers = {
|
||||
'Content-Type': 'application/json',
|
||||
**({'Authorization': f'Bearer {key}'} if key else {}),
|
||||
}
|
||||
|
||||
if ENABLE_FORWARD_USER_INFO_HEADERS and user:
|
||||
headers = include_user_info_headers(headers, user, request=request)
|
||||
if metadata and metadata.get('chat_id'):
|
||||
headers[FORWARD_SESSION_INFO_HEADER_CHAT_ID] = metadata.get('chat_id')
|
||||
|
||||
# Custom per-connection headers last so admin-set headers take precedence.
|
||||
if api_config and api_config.get('headers'):
|
||||
headers.update(await get_custom_headers(api_config['headers'], user, metadata, request=request))
|
||||
headers, cookies = await get_headers_and_cookies(request, url, key, api_config, metadata, user=user)
|
||||
|
||||
r = await session.request(
|
||||
method,
|
||||
url,
|
||||
data=payload,
|
||||
headers=headers,
|
||||
cookies=cookies,
|
||||
ssl=AIOHTTP_CLIENT_SESSION_SSL,
|
||||
timeout=get_client_timeout(stream=stream),
|
||||
)
|
||||
|
|
@ -249,24 +233,26 @@ class ConnectionVerificationForm(BaseModel):
|
|||
url: str
|
||||
key: str | None = None
|
||||
|
||||
config: dict | None = None
|
||||
|
||||
|
||||
@router.post('/verify')
|
||||
async def verify_connection(
|
||||
request: Request,
|
||||
form_data: ConnectionVerificationForm,
|
||||
user=Depends(get_admin_user),
|
||||
):
|
||||
"""Verify that an Ollama backend at *form_data.url* is reachable."""
|
||||
try:
|
||||
session = await get_session()
|
||||
headers: dict = {}
|
||||
if form_data.key:
|
||||
headers['Authorization'] = f'Bearer {form_data.key}'
|
||||
if ENABLE_FORWARD_USER_INFO_HEADERS and user:
|
||||
headers = include_user_info_headers(headers, user)
|
||||
headers, cookies = await get_headers_and_cookies(
|
||||
request, form_data.url, form_data.key, form_data.config, user=user
|
||||
)
|
||||
|
||||
async with session.get(
|
||||
f'{form_data.url}/api/version',
|
||||
headers=headers,
|
||||
cookies=cookies,
|
||||
ssl=AIOHTTP_CLIENT_SESSION_SSL,
|
||||
timeout=_MODEL_LIST_TIMEOUT,
|
||||
) as r:
|
||||
|
|
@ -403,9 +389,11 @@ async def get_all_models(request: Request, user: UserModel | None = None):
|
|||
for idx, url in enumerate(base_urls):
|
||||
api_config = resolve_api_config(api_configs, idx, url)
|
||||
if not api_config:
|
||||
tasks.append(send_get_request(f'{url}/api/tags', user=user))
|
||||
tasks.append(send_get_request(request, f'{url}/api/tags', user=user))
|
||||
elif api_config.get('enable', True):
|
||||
tasks.append(send_get_request(f'{url}/api/tags', api_config.get('key'), user=user))
|
||||
tasks.append(
|
||||
send_get_request(request, f'{url}/api/tags', api_config.get('key'), user=user, config=api_config)
|
||||
)
|
||||
else:
|
||||
tasks.append(asyncio.ensure_future(asyncio.sleep(0, None)))
|
||||
|
||||
|
|
@ -524,9 +512,11 @@ async def get_ollama_loaded_models(
|
|||
continue
|
||||
api_config = resolve_api_config(api_configs, idx, url)
|
||||
if not api_config:
|
||||
tasks.append(send_get_request(f'{url}/api/ps', user=user))
|
||||
tasks.append(send_get_request(request, f'{url}/api/ps', user=user))
|
||||
elif api_config.get('enable', True):
|
||||
tasks.append(send_get_request(f'{url}/api/ps', api_config.get('key'), user=user))
|
||||
tasks.append(
|
||||
send_get_request(request, f'{url}/api/ps', api_config.get('key'), user=user, config=api_config)
|
||||
)
|
||||
else:
|
||||
tasks.append(asyncio.ensure_future(asyncio.sleep(0, None)))
|
||||
|
||||
|
|
@ -570,7 +560,9 @@ async def get_ollama_versions(
|
|||
(await Config.get('ollama.api_configs', {})).get(url, {}),
|
||||
)
|
||||
if api_config.get('enable', True):
|
||||
tasks.append(send_get_request(f'{url}/api/version', api_config.get('key')))
|
||||
tasks.append(
|
||||
send_get_request(request, f'{url}/api/version', api_config.get('key'), user=user, config=api_config)
|
||||
)
|
||||
|
||||
raw = await asyncio.gather(*tasks)
|
||||
valid = [r for r in raw if r is not None]
|
||||
|
|
|
|||
|
|
@ -10,7 +10,6 @@ from urllib.parse import quote, urlparse
|
|||
import aiofiles
|
||||
import aiohttp
|
||||
from aiocache import cached
|
||||
from azure.identity import DefaultAzureCredential, get_bearer_token_provider
|
||||
from fastapi import APIRouter, Depends, HTTPException, Request, status
|
||||
from fastapi.responses import (
|
||||
FileResponse,
|
||||
|
|
@ -28,7 +27,6 @@ from open_webui.env import (
|
|||
BYPASS_MODEL_ACCESS_CONTROL,
|
||||
ENABLE_FORWARD_USER_INFO_HEADERS,
|
||||
ENABLE_OPENAI_API_PASSTHROUGH,
|
||||
FORWARD_SESSION_INFO_HEADER_CHAT_ID,
|
||||
MODELS_CACHE_TTL,
|
||||
REDIS_KEY_PREFIX,
|
||||
)
|
||||
|
|
@ -42,7 +40,7 @@ from open_webui.models.users import UserModel
|
|||
from open_webui.utils.access_control import check_model_access, has_connection_access, has_permission
|
||||
from open_webui.utils.anthropic import ANTHROPIC_VERSION, get_anthropic_models, is_anthropic_url
|
||||
from open_webui.utils.auth import get_admin_user, get_verified_user
|
||||
from open_webui.utils.headers import get_custom_headers, include_user_info_headers
|
||||
from open_webui.utils.headers import get_headers_and_cookies, include_user_info_headers
|
||||
from open_webui.utils.json_codec import JSONCodec
|
||||
from open_webui.utils.misc import convert_logit_bias_input_to_json
|
||||
from open_webui.utils.model_ids import strip_provider_model_prefix
|
||||
|
|
@ -152,87 +150,6 @@ def openai_reasoning_model_handler(payload):
|
|||
return payload
|
||||
|
||||
|
||||
async def get_headers_and_cookies(
|
||||
request: Request,
|
||||
url,
|
||||
key=None,
|
||||
config=None,
|
||||
metadata: dict | None = None,
|
||||
user: UserModel = None,
|
||||
):
|
||||
cookies = getattr(request, 'cookies', {}) if config.get('forward_cookies', False) else {}
|
||||
headers = {
|
||||
'Content-Type': 'application/json',
|
||||
**(
|
||||
{
|
||||
# LICENSE covers this Open WebUI upstream metadata identifier.
|
||||
# Do not alter, remove, obscure, or replace it except as LICENSE permits:
|
||||
# https://docs.openwebui.com/license.
|
||||
'HTTP-Referer': 'https://openwebui.com/',
|
||||
'X-Title': 'Open WebUI',
|
||||
}
|
||||
if 'openrouter.ai' in url
|
||||
else {}
|
||||
),
|
||||
}
|
||||
|
||||
if ENABLE_FORWARD_USER_INFO_HEADERS and user:
|
||||
headers = include_user_info_headers(headers, user, request=request)
|
||||
if metadata and metadata.get('chat_id'):
|
||||
headers[FORWARD_SESSION_INFO_HEADER_CHAT_ID] = metadata.get('chat_id')
|
||||
|
||||
token = None
|
||||
auth_type = config.get('auth_type')
|
||||
|
||||
if auth_type == 'bearer' or auth_type is None:
|
||||
# Default to bearer if not specified
|
||||
token = f'{key}'
|
||||
elif auth_type == 'none':
|
||||
token = None
|
||||
elif auth_type == 'session':
|
||||
token = request.state.token.credentials
|
||||
elif auth_type == 'system_oauth':
|
||||
oauth_token = None
|
||||
try:
|
||||
if request.cookies.get('oauth_session_id', None):
|
||||
oauth_token = await request.app.state.oauth_manager.get_oauth_token(
|
||||
user.id,
|
||||
request.cookies.get('oauth_session_id', None),
|
||||
)
|
||||
except Exception as e:
|
||||
log.error(f'Error getting OAuth token: {e}')
|
||||
|
||||
if oauth_token:
|
||||
token = f'{oauth_token.get("access_token", "")}'
|
||||
|
||||
elif auth_type in ('azure_ad', 'microsoft_entra_id'):
|
||||
token = get_microsoft_entra_id_access_token()
|
||||
|
||||
if token:
|
||||
headers['Authorization'] = f'Bearer {token}'
|
||||
|
||||
if config.get('headers') and isinstance(config.get('headers'), dict):
|
||||
custom_headers = await get_custom_headers(config.get('headers'), user, metadata, request=request)
|
||||
headers.update(custom_headers)
|
||||
|
||||
return headers, cookies
|
||||
|
||||
|
||||
def get_microsoft_entra_id_access_token():
|
||||
"""
|
||||
Get Microsoft Entra ID access token using DefaultAzureCredential for Azure OpenAI.
|
||||
Returns the token string or None if authentication fails.
|
||||
"""
|
||||
try:
|
||||
token_provider = get_bearer_token_provider(
|
||||
DefaultAzureCredential(), 'https://cognitiveservices.azure.com/.default'
|
||||
)
|
||||
return token_provider()
|
||||
except Exception as e:
|
||||
log.error(f'Error getting Microsoft Entra ID access token: {e}')
|
||||
return None
|
||||
|
||||
|
||||
##########################################
|
||||
#
|
||||
# API routes
|
||||
|
|
|
|||
|
|
@ -5,7 +5,10 @@ from typing import Any, Optional
|
|||
from urllib.parse import quote
|
||||
|
||||
import jwt
|
||||
from fastapi import Request
|
||||
from open_webui.env import (
|
||||
ENABLE_FORWARD_USER_INFO_HEADERS,
|
||||
FORWARD_SESSION_INFO_HEADER_CHAT_ID,
|
||||
FORWARD_USER_INFO_HEADER_AUTH_TYPE,
|
||||
FORWARD_USER_INFO_HEADER_JWT,
|
||||
FORWARD_USER_INFO_HEADER_JWT_EXPIRES_SECONDS,
|
||||
|
|
@ -161,3 +164,87 @@ def parse_custom_headers(
|
|||
parsed_headers[key] = quote(value, safe=punctuation + ' \t')
|
||||
|
||||
return parsed_headers
|
||||
|
||||
|
||||
async def get_headers_and_cookies(
|
||||
request: Request,
|
||||
url,
|
||||
key=None,
|
||||
config=None,
|
||||
metadata: dict | None = None,
|
||||
user=None,
|
||||
):
|
||||
config = config or {}
|
||||
cookies = getattr(request, 'cookies', {}) if config.get('forward_cookies', False) else {}
|
||||
headers = {
|
||||
'Content-Type': 'application/json',
|
||||
**(
|
||||
{
|
||||
# LICENSE covers this Open WebUI upstream metadata identifier.
|
||||
# Do not alter, remove, obscure, or replace it except as LICENSE permits:
|
||||
# https://docs.openwebui.com/license.
|
||||
'HTTP-Referer': 'https://openwebui.com/',
|
||||
'X-Title': 'Open WebUI',
|
||||
}
|
||||
if 'openrouter.ai' in url
|
||||
else {}
|
||||
),
|
||||
}
|
||||
|
||||
if ENABLE_FORWARD_USER_INFO_HEADERS and user:
|
||||
headers = include_user_info_headers(headers, user, request=request)
|
||||
if metadata and metadata.get('chat_id'):
|
||||
headers[FORWARD_SESSION_INFO_HEADER_CHAT_ID] = metadata.get('chat_id')
|
||||
|
||||
token = None
|
||||
auth_type = config.get('auth_type')
|
||||
|
||||
if auth_type == 'bearer' or auth_type is None:
|
||||
# Default to bearer if not specified
|
||||
token = key
|
||||
elif auth_type == 'none':
|
||||
token = None
|
||||
elif auth_type == 'session':
|
||||
token = request.state.token.credentials
|
||||
elif auth_type == 'system_oauth':
|
||||
oauth_token = None
|
||||
try:
|
||||
if request.cookies.get('oauth_session_id', None):
|
||||
oauth_token = await request.app.state.oauth_manager.get_oauth_token(
|
||||
user.id,
|
||||
request.cookies.get('oauth_session_id', None),
|
||||
)
|
||||
except Exception as e:
|
||||
log.error(f'Error getting OAuth token: {e}')
|
||||
|
||||
if oauth_token:
|
||||
token = f'{oauth_token.get("access_token", "")}'
|
||||
|
||||
elif auth_type in ('azure_ad', 'microsoft_entra_id'):
|
||||
token = get_microsoft_entra_id_access_token()
|
||||
|
||||
if token:
|
||||
headers['Authorization'] = f'Bearer {token}'
|
||||
|
||||
if config.get('headers') and isinstance(config.get('headers'), dict):
|
||||
custom_headers = await get_custom_headers(config.get('headers'), user, metadata, request=request)
|
||||
headers.update(custom_headers)
|
||||
|
||||
return headers, cookies
|
||||
|
||||
|
||||
def get_microsoft_entra_id_access_token():
|
||||
"""
|
||||
Get Microsoft Entra ID access token using DefaultAzureCredential for Azure OpenAI.
|
||||
Returns the token string or None if authentication fails.
|
||||
"""
|
||||
from azure.identity import DefaultAzureCredential, get_bearer_token_provider
|
||||
|
||||
try:
|
||||
token_provider = get_bearer_token_provider(
|
||||
DefaultAzureCredential(), 'https://cognitiveservices.azure.com/.default'
|
||||
)
|
||||
return token_provider()
|
||||
except Exception as e:
|
||||
log.error(f'Error getting Microsoft Entra ID access token: {e}')
|
||||
return None
|
||||
|
|
|
|||
|
|
@ -77,9 +77,29 @@
|
|||
// remove trailing slash from url
|
||||
url = url.replace(/\/$/, '');
|
||||
|
||||
let _headers = null;
|
||||
|
||||
if (headers) {
|
||||
try {
|
||||
_headers = JSON.parse(headers);
|
||||
if (typeof _headers !== 'object' || Array.isArray(_headers)) {
|
||||
_headers = null;
|
||||
throw new Error('Headers must be a valid JSON object');
|
||||
}
|
||||
headers = JSON.stringify(_headers, null, 2);
|
||||
} catch (error) {
|
||||
toast.error($i18n.t('Headers must be a valid JSON object'));
|
||||
return;
|
||||
}
|
||||
}
|
||||
|
||||
const res = await verifyOllamaConnection(localStorage.token, {
|
||||
url,
|
||||
key
|
||||
key,
|
||||
config: {
|
||||
auth_type,
|
||||
...(_headers ? { headers: _headers } : {})
|
||||
}
|
||||
}).catch((error) => {
|
||||
toast.error(`${error}`);
|
||||
});
|
||||
|
|
|
|||
Loading…
Add table
Reference in a new issue