refac: ollama

This commit is contained in:
Timothy Jaeryang Baek
2026-04-08 14:32:34 -07:00
parent a775fc9b50
commit e51b661af0
+155 -353
View File
@@ -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)}