refac: ollama
This commit is contained in:
@@ -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)}
|
||||
|
||||
Reference in New Issue
Block a user