From e51b661af0e71a24f041428f328fcc6e97a15262 Mon Sep 17 00:00:00 2001 From: Timothy Jaeryang Baek Date: Wed, 8 Apr 2026 14:32:34 -0700 Subject: [PATCH] refac: ollama --- backend/open_webui/routers/ollama.py | 508 ++++++++------------------- 1 file changed, 155 insertions(+), 353 deletions(-) diff --git a/backend/open_webui/routers/ollama.py b/backend/open_webui/routers/ollama.py index fd2a088a6..93745440c 100644 --- a/backend/open_webui/routers/ollama.py +++ b/backend/open_webui/routers/ollama.py @@ -15,7 +15,7 @@ from typing import Optional, Union from urllib.parse import urlparse import aiohttp from aiocache import cached -import requests + from open_webui.utils.headers import include_user_info_headers from open_webui.models.chats import Chats @@ -107,19 +107,24 @@ async def send_get_request(url, key=None, user: UserModel = None): return None -async def send_post_request( +async def send_request( url: str, - payload: Union[str, bytes], - stream: bool = True, + method: str = 'POST', + *, + payload: Optional[Union[str, bytes]] = None, key: Optional[str] = None, - content_type: Optional[str] = None, user: UserModel = None, + stream: bool = False, + content_type: Optional[str] = None, metadata: Optional[dict] = None, ): r = None streaming = False try: - session = aiohttp.ClientSession(trust_env=True, timeout=aiohttp.ClientTimeout(total=AIOHTTP_CLIENT_TIMEOUT)) + session = aiohttp.ClientSession( + trust_env=True, + timeout=aiohttp.ClientTimeout(total=AIOHTTP_CLIENT_TIMEOUT), + ) headers = { 'Content-Type': 'application/json', @@ -131,32 +136,29 @@ async def send_post_request( if metadata and metadata.get('chat_id'): headers[FORWARD_SESSION_INFO_HEADER_CHAT_ID] = metadata.get('chat_id') - r = await session.post( - url, - data=payload, - headers=headers, + r = await session.request( + method, url, data=payload, headers=headers, ssl=AIOHTTP_CLIENT_SESSION_SSL, ) - if r.ok is False: + if not r.ok: try: res = await r.json() - await cleanup_response(r, session) if 'error' in res: raise HTTPException(status_code=r.status, detail=res['error']) - except HTTPException as e: - raise e # Re-raise HTTPException to be handled by FastAPI + except HTTPException: + raise except Exception as e: log.error(f'Failed to parse error response: {e}') - raise HTTPException( - status_code=r.status, - detail=f'Open WebUI: Server Connection Error', - ) + raise HTTPException( + status_code=r.status, + detail='Open WebUI: Server Connection Error', + ) + + r.raise_for_status() - r.raise_for_status() # Raises an error for bad responses (4xx, 5xx) if stream: response_headers = dict(r.headers) - if content_type: response_headers['Content-Type'] = content_type @@ -167,17 +169,17 @@ async def send_post_request( headers=response_headers, ) else: - res = await r.json() - return res + try: + return await r.json() + except Exception: + return None - except HTTPException as e: - raise e # Re-raise HTTPException to be handled by FastAPI + except HTTPException: + raise except Exception as e: - detail = f'Ollama: {e}' - raise HTTPException( status_code=r.status if r else 500, - detail=detail if e else 'Open WebUI: Server Connection Error', + detail=f'Ollama: {e}' if str(e) else 'Open WebUI: Server Connection Error', ) finally: if not streaming: @@ -430,40 +432,7 @@ async def get_ollama_tags(request: Request, url_idx: Optional[int] = None, user= else: url = request.app.state.config.OLLAMA_BASE_URLS[url_idx] key = get_api_key(url_idx, url, request.app.state.config.OLLAMA_API_CONFIGS) - - r = None - try: - headers = { - **({'Authorization': f'Bearer {key}'} if key else {}), - } - - if ENABLE_FORWARD_USER_INFO_HEADERS and user: - headers = include_user_info_headers(headers, user) - - r = requests.request( - method='GET', - url=f'{url}/api/tags', - headers=headers, - ) - r.raise_for_status() - - models = r.json() - except Exception as e: - log.exception(e) - - detail = None - if r is not None: - try: - res = r.json() - if 'error' in res: - detail = f'Ollama: {res["error"]}' - except Exception: - detail = f'Ollama: {e}' - - raise HTTPException( - status_code=r.status_code if r else 500, - detail=detail if detail else 'Open WebUI: Server Connection Error', - ) + models = await send_request(f'{url}/api/tags', 'GET', key=key, user=user) if user.role == 'user' and not BYPASS_MODEL_ACCESS_CONTROL: models['models'] = await get_filtered_models(models, user) @@ -569,29 +538,7 @@ async def get_ollama_versions(request: Request, url_idx: Optional[int] = None): ) else: url = request.app.state.config.OLLAMA_BASE_URLS[url_idx] - - r = None - try: - r = requests.request(method='GET', url=f'{url}/api/version') - r.raise_for_status() - - return r.json() - except Exception as e: - log.exception(e) - - detail = None - if r is not None: - try: - res = r.json() - if 'error' in res: - detail = f'Ollama: {res["error"]}' - except Exception: - detail = f'Ollama: {e}' - - raise HTTPException( - status_code=r.status_code if r else 500, - detail=detail if detail else 'Open WebUI: Server Connection Error', - ) + return await send_request(f'{url}/api/version', 'GET') else: return {'version': False} @@ -640,10 +587,9 @@ async def unload_model( payload = {'model': model_name, 'keep_alive': 0, 'prompt': ''} try: - res = await send_post_request( - url=f'{url}/api/generate', + res = await send_request( + f'{url}/api/generate', payload=json.dumps(payload), - stream=False, key=key, user=user, ) @@ -681,11 +627,12 @@ async def pull_model( # Admin should be able to pull models from any source payload = {**form_data, 'insecure': True} - return await send_post_request( - url=f'{url}/api/pull', + return await send_request( + f'{url}/api/pull', payload=json.dumps(payload), key=get_api_key(url_idx, url, request.app.state.config.OLLAMA_API_CONFIGS), user=user, + stream=True, ) @@ -721,11 +668,12 @@ async def push_model( url = request.app.state.config.OLLAMA_BASE_URLS[url_idx] log.debug(f'url: {url}') - return await send_post_request( - url=f'{url}/api/push', + return await send_request( + f'{url}/api/push', payload=form_data.model_dump_json(exclude_none=True).encode(), key=get_api_key(url_idx, url, request.app.state.config.OLLAMA_API_CONFIGS), user=user, + stream=True, ) @@ -751,11 +699,12 @@ async def create_model( log.debug(f'form_data: {form_data}') url = request.app.state.config.OLLAMA_BASE_URLS[url_idx] - return await send_post_request( - url=f'{url}/api/create', + return await send_request( + f'{url}/api/create', payload=form_data.model_dump_json(exclude_none=True).encode(), key=get_api_key(url_idx, url, request.app.state.config.OLLAMA_API_CONFIGS), user=user, + stream=True, ) @@ -790,41 +739,13 @@ async def copy_model( url = request.app.state.config.OLLAMA_BASE_URLS[url_idx] key = get_api_key(url_idx, url, request.app.state.config.OLLAMA_API_CONFIGS) - try: - 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) - - r = requests.request( - method='POST', - url=f'{url}/api/copy', - headers=headers, - data=form_data.model_dump_json(exclude_none=True).encode(), - ) - r.raise_for_status() - - log.debug(f'r.text: {r.text}') - return True - except Exception as e: - log.exception(e) - - detail = None - if r is not None: - try: - res = r.json() - if 'error' in res: - detail = f'Ollama: {res["error"]}' - except Exception: - detail = f'Ollama: {e}' - - raise HTTPException( - status_code=r.status_code if r else 500, - detail=detail if detail else 'Open WebUI: Server Connection Error', - ) + await send_request( + f'{url}/api/copy', + payload=form_data.model_dump_json(exclude_none=True).encode(), + key=key, + user=user, + ) + return True @router.delete('/api/delete') @@ -858,42 +779,13 @@ async def delete_model( url = request.app.state.config.OLLAMA_BASE_URLS[url_idx] key = get_api_key(url_idx, url, request.app.state.config.OLLAMA_API_CONFIGS) - r = None - try: - 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) - - r = requests.request( - method='DELETE', - url=f'{url}/api/delete', - headers=headers, - json=form_data, - ) - r.raise_for_status() - - log.debug(f'r.text: {r.text}') - return True - except Exception as e: - log.exception(e) - - detail = None - if r is not None: - try: - res = r.json() - if 'error' in res: - detail = f'Ollama: {res["error"]}' - except Exception: - detail = f'Ollama: {e}' - - raise HTTPException( - status_code=r.status_code if r else 500, - detail=detail if detail else 'Open WebUI: Server Connection Error', - ) + await send_request( + f'{url}/api/delete', 'DELETE', + payload=json.dumps(form_data), + key=key, + user=user, + ) + return True @router.post('/api/show') @@ -920,35 +812,12 @@ async def show_model_info(request: Request, form_data: ModelNameForm, user=Depen url = request.app.state.config.OLLAMA_BASE_URLS[url_idx] key = get_api_key(url_idx, url, request.app.state.config.OLLAMA_API_CONFIGS) - try: - 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) - - r = requests.request(method='POST', url=f'{url}/api/show', headers=headers, json=form_data) - r.raise_for_status() - - return r.json() - except Exception as e: - log.exception(e) - - detail = None - if r is not None: - try: - res = r.json() - if 'error' in res: - detail = f'Ollama: {res["error"]}' - except Exception: - detail = f'Ollama: {e}' - - raise HTTPException( - status_code=r.status_code if r else 500, - detail=detail if detail else 'Open WebUI: Server Connection Error', - ) + return await send_request( + f'{url}/api/show', + payload=json.dumps(form_data), + key=key, + user=user, + ) class GenerateEmbedForm(BaseModel): @@ -1004,41 +873,12 @@ async def embed( if prefix_id: form_data.model = form_data.model.replace(f'{prefix_id}.', '') - try: - 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) - - r = requests.request( - method='POST', - url=f'{url}/api/embed', - headers=headers, - data=form_data.model_dump_json(exclude_none=True).encode(), - ) - r.raise_for_status() - - data = r.json() - return data - except Exception as e: - log.exception(e) - - detail = None - if r is not None: - try: - res = r.json() - if 'error' in res: - detail = f'Ollama: {res["error"]}' - except Exception: - detail = f'Ollama: {e}' - - raise HTTPException( - status_code=r.status_code if r else 500, - detail=detail if detail else 'Open WebUI: Server Connection Error', - ) + return await send_request( + f'{url}/api/embed', + payload=form_data.model_dump_json(exclude_none=True).encode(), + key=key, + user=user, + ) class GenerateEmbeddingsForm(BaseModel): @@ -1089,41 +929,12 @@ async def embeddings( if prefix_id: form_data.model = form_data.model.replace(f'{prefix_id}.', '') - try: - 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) - - r = requests.request( - method='POST', - url=f'{url}/api/embeddings', - headers=headers, - data=form_data.model_dump_json(exclude_none=True).encode(), - ) - r.raise_for_status() - - data = r.json() - return data - except Exception as e: - log.exception(e) - - detail = None - if r is not None: - try: - res = r.json() - if 'error' in res: - detail = f'Ollama: {res["error"]}' - except Exception: - detail = f'Ollama: {e}' - - raise HTTPException( - status_code=r.status_code if r else 500, - detail=detail if detail else 'Open WebUI: Server Connection Error', - ) + return await send_request( + f'{url}/api/embeddings', + payload=form_data.model_dump_json(exclude_none=True).encode(), + key=key, + user=user, + ) class GenerateCompletionForm(BaseModel): @@ -1175,11 +986,12 @@ async def generate_completion( if prefix_id: form_data.model = form_data.model.replace(f'{prefix_id}.', '') - return await send_post_request( - url=f'{url}/api/generate', + return await send_request( + f'{url}/api/generate', payload=form_data.model_dump_json(exclude_none=True).encode(), key=get_api_key(url_idx, url, request.app.state.config.OLLAMA_API_CONFIGS), user=user, + stream=True, ) @@ -1319,13 +1131,13 @@ async def generate_chat_completion( if prefix_id: payload['model'] = payload['model'].replace(f'{prefix_id}.', '') - return await send_post_request( - url=f'{url}/api/chat', + return await send_request( + f'{url}/api/chat', payload=json.dumps(payload), - stream=form_data.stream, key=get_api_key(url_idx, url, request.app.state.config.OLLAMA_API_CONFIGS), - content_type='application/x-ndjson', user=user, + stream=form_data.stream, + content_type='application/x-ndjson', metadata=metadata, ) @@ -1429,12 +1241,12 @@ async def generate_openai_completion( if prefix_id: payload['model'] = payload['model'].replace(f'{prefix_id}.', '') - return await send_post_request( - url=f'{url}/v1/completions', + return await send_request( + f'{url}/v1/completions', payload=json.dumps(payload), - stream=payload.get('stream', False), key=get_api_key(url_idx, url, request.app.state.config.OLLAMA_API_CONFIGS), user=user, + stream=payload.get('stream', False), metadata=metadata, ) @@ -1514,12 +1326,12 @@ async def generate_openai_chat_completion( if prefix_id: payload['model'] = payload['model'].replace(f'{prefix_id}.', '') - return await send_post_request( - url=f'{url}/v1/chat/completions', + return await send_request( + f'{url}/v1/chat/completions', payload=json.dumps(payload), - stream=payload.get('stream', False), key=get_api_key(url_idx, url, request.app.state.config.OLLAMA_API_CONFIGS), user=user, + stream=payload.get('stream', False), metadata=metadata, ) @@ -1586,13 +1398,13 @@ async def generate_anthropic_messages( if prefix_id: payload['model'] = payload['model'].replace(f'{prefix_id}.', '') - return await send_post_request( - url=f'{url}/v1/messages', + return await send_request( + f'{url}/v1/messages', payload=json.dumps(payload), - stream=payload.get('stream', False), - content_type='text/event-stream' if payload.get('stream', False) else None, key=get_api_key(url_idx, url, request.app.state.config.OLLAMA_API_CONFIGS), user=user, + stream=payload.get('stream', False), + content_type='text/event-stream' if payload.get('stream', False) else None, ) @@ -1664,13 +1476,13 @@ async def generate_responses( if prefix_id: payload['model'] = payload['model'].replace(f'{prefix_id}.', '') - return await send_post_request( - url=f'{url}/v1/responses', + return await send_request( + f'{url}/v1/responses', payload=json.dumps(payload), - stream=payload.get('stream', False), - content_type='text/event-stream' if payload.get('stream', False) else None, key=get_api_key(url_idx, url, request.app.state.config.OLLAMA_API_CONFIGS), user=user, + stream=payload.get('stream', False), + content_type='text/event-stream' if payload.get('stream', False) else None, ) @@ -1697,36 +1509,17 @@ async def get_openai_models( else: url = request.app.state.config.OLLAMA_BASE_URLS[url_idx] - try: - r = requests.request(method='GET', url=f'{url}/api/tags') - r.raise_for_status() + model_list = await send_request(f'{url}/api/tags', 'GET') - model_list = r.json() - - models = [ - { - 'id': model['model'], - 'object': 'model', - 'created': int(time.time()), - 'owned_by': 'openai', - } - for model in models['models'] - ] - except Exception as e: - log.exception(e) - error_detail = 'Open WebUI: Server Connection Error' - if r is not None: - try: - res = r.json() - if 'error' in res: - error_detail = f'Ollama: {res["error"]}' - except Exception: - error_detail = f'Ollama: {e}' - - raise HTTPException( - status_code=r.status_code if r else 500, - detail=error_detail, - ) + models = [ + { + 'id': model['model'], + 'object': 'model', + 'created': int(time.time()), + 'owned_by': 'openai', + } + for model in model_list.get('models', []) + ] if user.role == 'user' and not BYPASS_MODEL_ACCESS_CONTROL: # Filter models based on user access control @@ -1812,13 +1605,16 @@ async def download_file_stream(ollama_url, file_url, file_path, file_name, chunk file.close() hashed = calculate_sha256(file_path, chunk_size) - with open(file_path, 'rb') as file: - chunk_size = 1024 * 1024 * 2 - url = f'{ollama_url}/api/blobs/sha256:{hashed}' - with requests.Session() as session: - response = session.post(url, data=file, timeout=30) + with open(file_path, 'rb') as f: + blob_data = f.read() - if response.ok: + url = f'{ollama_url}/api/blobs/sha256:{hashed}' + blob_timeout = aiohttp.ClientTimeout(total=30) + async with aiohttp.ClientSession(timeout=blob_timeout, trust_env=True) as blob_session: + async with blob_session.post( + url, data=blob_data, ssl=AIOHTTP_CLIENT_SESSION_SSL + ) as blob_response: + if blob_response.ok: res = { 'done': done, 'blob': f'sha256:{hashed}', @@ -1914,47 +1710,53 @@ async def upload_model( # --- P3: Upload to ollama /api/blobs --- with open(file_path, 'rb') as f: - url = f'{ollama_url}/api/blobs/sha256:{file_hash}' - response = requests.post(url, data=f) + blob_data = f.read() - if response.ok: - log.info(f'Uploaded to /api/blobs') # DEBUG - # Remove local file - os.remove(file_path) + url = f'{ollama_url}/api/blobs/sha256:{file_hash}' + upload_timeout = aiohttp.ClientTimeout(total=AIOHTTP_CLIENT_TIMEOUT) + async with aiohttp.ClientSession(timeout=upload_timeout, trust_env=True) as upload_session: + async with upload_session.post( + url, data=blob_data, ssl=AIOHTTP_CLIENT_SESSION_SSL + ) as response: + if not response.ok: + raise Exception('Ollama: Could not create blob, Please try again.') - # Create model in ollama - model_name, ext = os.path.splitext(filename) - log.info(f'Created Model: {model_name}') # DEBUG + log.info(f'Uploaded to /api/blobs') # DEBUG + # Remove local file + os.remove(file_path) - create_payload = { - 'model': model_name, - # Reference the file by its original name => the uploaded blob's digest - 'files': {filename: f'sha256:{file_hash}'}, - } - log.info(f'Model Payload: {create_payload}') # DEBUG + # Create model in ollama + model_name, ext = os.path.splitext(filename) + log.info(f'Created Model: {model_name}') # DEBUG - # Call ollama /api/create - # https://github.com/ollama/ollama/blob/main/docs/api.md#create-a-model - create_resp = requests.post( - url=f'{ollama_url}/api/create', + create_payload = { + 'model': model_name, + # Reference the file by its original name => the uploaded blob's digest + 'files': {filename: f'sha256:{file_hash}'}, + } + log.info(f'Model Payload: {create_payload}') # DEBUG + + # Call ollama /api/create + # https://github.com/ollama/ollama/blob/main/docs/api.md#create-a-model + async with aiohttp.ClientSession(timeout=upload_timeout, trust_env=True) as create_session: + async with create_session.post( + f'{ollama_url}/api/create', headers={'Content-Type': 'application/json'}, data=json.dumps(create_payload), - ) - - if create_resp.ok: - log.info(f'API SUCCESS!') # DEBUG - done_msg = { - 'done': True, - 'blob': f'sha256:{file_hash}', - 'name': filename, - 'model_created': model_name, - } - yield f'data: {json.dumps(done_msg)}\n\n' - else: - raise Exception(f'Failed to create model in Ollama. {create_resp.text}') - - else: - raise Exception('Ollama: Could not create blob, Please try again.') + ssl=AIOHTTP_CLIENT_SESSION_SSL, + ) as create_resp: + if create_resp.ok: + log.info(f'API SUCCESS!') # DEBUG + done_msg = { + 'done': True, + 'blob': f'sha256:{file_hash}', + 'name': filename, + 'model_created': model_name, + } + yield f'data: {json.dumps(done_msg)}\n\n' + else: + resp_text = await create_resp.text() + raise Exception(f'Failed to create model in Ollama. {resp_text}') except Exception as e: res = {'error': str(e)}