chore: format
This commit is contained in:
@@ -88,25 +88,29 @@ async def get_user_analytics(
|
||||
token_usage = ChatMessages.get_token_usage_by_user(
|
||||
start_date=start_date, end_date=end_date, db=db
|
||||
)
|
||||
|
||||
|
||||
# Get user info for top users
|
||||
top_user_ids = [uid for uid, _ in sorted(counts.items(), key=lambda x: -x[1])[:limit]]
|
||||
top_user_ids = [
|
||||
uid for uid, _ in sorted(counts.items(), key=lambda x: -x[1])[:limit]
|
||||
]
|
||||
user_info = {u.id: u for u in Users.get_users_by_user_ids(top_user_ids, db=db)}
|
||||
|
||||
|
||||
users = []
|
||||
for user_id in top_user_ids:
|
||||
u = user_info.get(user_id)
|
||||
tokens = token_usage.get(user_id, {})
|
||||
users.append(UserAnalyticsEntry(
|
||||
user_id=user_id,
|
||||
name=u.name if u else None,
|
||||
email=u.email if u else None,
|
||||
count=counts[user_id],
|
||||
input_tokens=tokens.get("input_tokens", 0),
|
||||
output_tokens=tokens.get("output_tokens", 0),
|
||||
total_tokens=tokens.get("total_tokens", 0),
|
||||
))
|
||||
|
||||
users.append(
|
||||
UserAnalyticsEntry(
|
||||
user_id=user_id,
|
||||
name=u.name if u else None,
|
||||
email=u.email if u else None,
|
||||
count=counts[user_id],
|
||||
input_tokens=tokens.get("input_tokens", 0),
|
||||
output_tokens=tokens.get("output_tokens", 0),
|
||||
total_tokens=tokens.get("total_tokens", 0),
|
||||
)
|
||||
)
|
||||
|
||||
return UserAnalyticsResponse(users=users)
|
||||
|
||||
|
||||
@@ -168,7 +172,7 @@ async def get_summary(
|
||||
chat_counts = ChatMessages.get_message_count_by_chat(
|
||||
start_date=start_date, end_date=end_date, group_id=group_id, db=db
|
||||
)
|
||||
|
||||
|
||||
return SummaryResponse(
|
||||
total_messages=sum(model_counts.values()),
|
||||
total_chats=len(chat_counts),
|
||||
@@ -317,9 +321,7 @@ async def get_model_chats(
|
||||
if isinstance(content, str):
|
||||
first_message = content[:200]
|
||||
elif isinstance(content, list):
|
||||
text_parts = [
|
||||
b.get("text", "") for b in content if isinstance(b, dict)
|
||||
]
|
||||
text_parts = [b.get("text", "") for b in content if isinstance(b, dict)]
|
||||
first_message = " ".join(text_parts)[:200]
|
||||
|
||||
# Get user info
|
||||
@@ -331,7 +333,6 @@ async def get_model_chats(
|
||||
# Timestamps from messages
|
||||
updated_at = max(m.created_at for m in messages) if messages else 0
|
||||
|
||||
|
||||
chats_data.append(
|
||||
ModelChatEntry(
|
||||
chat_id=chat_id,
|
||||
@@ -387,24 +388,24 @@ async def get_model_overview(
|
||||
|
||||
# Get feedback history per day
|
||||
history_counts: dict[str, dict] = defaultdict(lambda: {"won": 0, "lost": 0})
|
||||
|
||||
|
||||
# Calculate start date for history
|
||||
now = datetime.now()
|
||||
start_dt = None
|
||||
if days > 0:
|
||||
start_dt = now - timedelta(days=days)
|
||||
|
||||
|
||||
for chat_id in chat_ids:
|
||||
feedbacks = Feedbacks.get_feedbacks_by_chat_id(chat_id, db=db)
|
||||
for fb in feedbacks:
|
||||
if fb.data and "rating" in fb.data:
|
||||
rating = fb.data["rating"]
|
||||
fb_date = datetime.fromtimestamp(fb.created_at)
|
||||
|
||||
|
||||
# Filter by date range
|
||||
if start_dt and fb_date < start_dt:
|
||||
continue
|
||||
|
||||
|
||||
date_str = fb_date.strftime("%Y-%m-%d")
|
||||
if rating == 1:
|
||||
history_counts[date_str]["won"] += 1
|
||||
@@ -423,15 +424,17 @@ async def get_model_overview(
|
||||
current = datetime.strptime(min_date, "%Y-%m-%d")
|
||||
else:
|
||||
current = now
|
||||
|
||||
|
||||
while current <= end_dt:
|
||||
date_str = current.strftime("%Y-%m-%d")
|
||||
counts = history_counts.get(date_str, {"won": 0, "lost": 0})
|
||||
history.append(HistoryEntry(
|
||||
date=date_str,
|
||||
won=counts["won"],
|
||||
lost=counts["lost"],
|
||||
))
|
||||
history.append(
|
||||
HistoryEntry(
|
||||
date=date_str,
|
||||
won=counts["won"],
|
||||
lost=counts["lost"],
|
||||
)
|
||||
)
|
||||
current += timedelta(days=1)
|
||||
|
||||
# Get chat tags
|
||||
|
||||
@@ -57,7 +57,6 @@ from open_webui.env import (
|
||||
ENABLE_FORWARD_USER_INFO_HEADERS,
|
||||
)
|
||||
|
||||
|
||||
router = APIRouter()
|
||||
|
||||
# Constants
|
||||
|
||||
@@ -98,7 +98,7 @@ def create_session_response(
|
||||
"""
|
||||
Create JWT token and build session response for a user.
|
||||
Shared helper for signin, signup, ldap_auth, add_user, and token_exchange endpoints.
|
||||
|
||||
|
||||
Args:
|
||||
request: FastAPI request object
|
||||
user: User object
|
||||
@@ -558,7 +558,9 @@ async def ldap_auth(
|
||||
except Exception as e:
|
||||
log.error(f"Failed to sync groups for user {user.id}: {e}")
|
||||
|
||||
return create_session_response(request, user, db, response, set_cookie=True)
|
||||
return create_session_response(
|
||||
request, user, db, response, set_cookie=True
|
||||
)
|
||||
else:
|
||||
raise HTTPException(400, detail=ERROR_MESSAGES.INVALID_CRED)
|
||||
else:
|
||||
|
||||
@@ -1633,9 +1633,7 @@ async def update_message_by_id(
|
||||
if (
|
||||
user.role != "admin"
|
||||
and message.user_id != user.id
|
||||
and not channel_has_access(
|
||||
user.id, channel, permission="read", db=db
|
||||
)
|
||||
and not channel_has_access(user.id, channel, permission="read", db=db)
|
||||
):
|
||||
raise HTTPException(
|
||||
status_code=status.HTTP_403_FORBIDDEN, detail=ERROR_MESSAGES.DEFAULT()
|
||||
|
||||
@@ -33,7 +33,6 @@ from fastapi.responses import FileResponse, StreamingResponse
|
||||
from open_webui.utils.auth import get_admin_user, get_verified_user
|
||||
from open_webui.utils.access_control import has_permission
|
||||
|
||||
|
||||
log = logging.getLogger(__name__)
|
||||
|
||||
|
||||
|
||||
@@ -29,7 +29,6 @@ from pydantic import BaseModel, HttpUrl
|
||||
from open_webui.internal.db import get_session
|
||||
from sqlalchemy.orm import Session
|
||||
|
||||
|
||||
log = logging.getLogger(__name__)
|
||||
|
||||
|
||||
|
||||
@@ -22,7 +22,6 @@ from sqlalchemy.orm import Session
|
||||
|
||||
from open_webui.utils.auth import get_admin_user, get_verified_user
|
||||
|
||||
|
||||
log = logging.getLogger(__name__)
|
||||
|
||||
router = APIRouter()
|
||||
|
||||
@@ -14,7 +14,11 @@ import requests
|
||||
from fastapi import APIRouter, Depends, HTTPException, Request, UploadFile
|
||||
from fastapi.responses import FileResponse
|
||||
|
||||
from open_webui.config import CACHE_DIR, IMAGE_AUTO_SIZE_MODELS_REGEX_PATTERN, IMAGE_URL_RESPONSE_MODELS_REGEX_PATTERN
|
||||
from open_webui.config import (
|
||||
CACHE_DIR,
|
||||
IMAGE_AUTO_SIZE_MODELS_REGEX_PATTERN,
|
||||
IMAGE_URL_RESPONSE_MODELS_REGEX_PATTERN,
|
||||
)
|
||||
from open_webui.constants import ERROR_MESSAGES
|
||||
from open_webui.retrieval.web.utils import validate_url
|
||||
from open_webui.env import ENABLE_FORWARD_USER_INFO_HEADERS
|
||||
@@ -199,9 +203,8 @@ async def update_config(
|
||||
|
||||
request.app.state.config.IMAGE_GENERATION_ENGINE = form_data.IMAGE_GENERATION_ENGINE
|
||||
set_image_model(request, form_data.IMAGE_GENERATION_MODEL)
|
||||
if (
|
||||
form_data.IMAGE_SIZE == "auto"
|
||||
and not re.match(IMAGE_AUTO_SIZE_MODELS_REGEX_PATTERN, form_data.IMAGE_GENERATION_MODEL)
|
||||
if form_data.IMAGE_SIZE == "auto" and not re.match(
|
||||
IMAGE_AUTO_SIZE_MODELS_REGEX_PATTERN, form_data.IMAGE_GENERATION_MODEL
|
||||
):
|
||||
raise HTTPException(
|
||||
status_code=400,
|
||||
@@ -610,7 +613,10 @@ async def image_generations(
|
||||
),
|
||||
**(
|
||||
{}
|
||||
if re.match(IMAGE_URL_RESPONSE_MODELS_REGEX_PATTERN, request.app.state.config.IMAGE_GENERATION_MODEL)
|
||||
if re.match(
|
||||
IMAGE_URL_RESPONSE_MODELS_REGEX_PATTERN,
|
||||
request.app.state.config.IMAGE_GENERATION_MODEL,
|
||||
)
|
||||
else {"response_format": "b64_json"}
|
||||
),
|
||||
**(
|
||||
@@ -912,7 +918,9 @@ async def image_edits(
|
||||
form_data.image = await load_url_image(form_data.image)
|
||||
elif isinstance(form_data.image, list):
|
||||
# Load all images in parallel for better performance
|
||||
form_data.image = list(await asyncio.gather(*[load_url_image(img) for img in form_data.image]))
|
||||
form_data.image = list(
|
||||
await asyncio.gather(*[load_url_image(img) for img in form_data.image])
|
||||
)
|
||||
except Exception as e:
|
||||
raise HTTPException(status_code=400, detail=ERROR_MESSAGES.DEFAULT(e))
|
||||
|
||||
@@ -947,7 +955,10 @@ async def image_edits(
|
||||
**({"size": size} if size else {}),
|
||||
**(
|
||||
{}
|
||||
if re.match(IMAGE_URL_RESPONSE_MODELS_REGEX_PATTERN, request.app.state.config.IMAGE_EDIT_MODEL)
|
||||
if re.match(
|
||||
IMAGE_URL_RESPONSE_MODELS_REGEX_PATTERN,
|
||||
request.app.state.config.IMAGE_EDIT_MODEL,
|
||||
)
|
||||
else {"response_format": "b64_json"}
|
||||
),
|
||||
}
|
||||
|
||||
@@ -36,7 +36,6 @@ from open_webui.models.access_grants import AccessGrants, has_public_read_access
|
||||
from open_webui.config import BYPASS_ADMIN_ACCESS_CONTROL
|
||||
from open_webui.models.models import Models, ModelForm
|
||||
|
||||
|
||||
log = logging.getLogger(__name__)
|
||||
|
||||
router = APIRouter()
|
||||
@@ -358,7 +357,7 @@ async def reindex_knowledge_base_metadata_embeddings(
|
||||
user=Depends(get_admin_user),
|
||||
):
|
||||
"""Batch embed all existing knowledge bases. Admin only.
|
||||
|
||||
|
||||
NOTE: We intentionally do NOT use Depends(get_session) here.
|
||||
This endpoint loops through ALL knowledge bases and calls embed_knowledge_base_metadata()
|
||||
for each one, making N external embedding API calls. Holding a session during
|
||||
@@ -540,9 +539,7 @@ async def update_knowledge_access_by_id(
|
||||
detail=ERROR_MESSAGES.ACCESS_PROHIBITED,
|
||||
)
|
||||
|
||||
AccessGrants.set_access_grants(
|
||||
"knowledge", id, form_data.access_grants, db=db
|
||||
)
|
||||
AccessGrants.set_access_grants("knowledge", id, form_data.access_grants, db=db)
|
||||
|
||||
return KnowledgeFilesResponse(
|
||||
**Knowledges.get_knowledge_by_id(id=id, db=db).model_dump(),
|
||||
|
||||
@@ -345,9 +345,7 @@ async def update_note_access_by_id(
|
||||
status_code=status.HTTP_403_FORBIDDEN, detail=ERROR_MESSAGES.DEFAULT()
|
||||
)
|
||||
|
||||
AccessGrants.set_access_grants(
|
||||
"note", id, form_data.access_grants, db=db
|
||||
)
|
||||
AccessGrants.set_access_grants("note", id, form_data.access_grants, db=db)
|
||||
|
||||
return Notes.get_note_by_id(id, db=db)
|
||||
|
||||
|
||||
@@ -54,7 +54,6 @@ from open_webui.utils.misc import (
|
||||
from open_webui.utils.auth import get_admin_user, get_verified_user
|
||||
from open_webui.utils.headers import include_user_info_headers
|
||||
|
||||
|
||||
log = logging.getLogger(__name__)
|
||||
|
||||
|
||||
@@ -805,67 +804,77 @@ def convert_to_azure_payload(url, payload: dict, api_version: str):
|
||||
def convert_to_responses_payload(payload: dict) -> dict:
|
||||
"""
|
||||
Convert Chat Completions payload to Responses API format.
|
||||
|
||||
|
||||
Chat Completions: { messages: [{role, content}], ... }
|
||||
Responses API: { input: [{type: "message", role, content: [...]}], instructions: "system" }
|
||||
"""
|
||||
messages = payload.pop("messages", [])
|
||||
|
||||
|
||||
system_content = ""
|
||||
input_items = []
|
||||
|
||||
|
||||
for msg in messages:
|
||||
role = msg.get("role", "user")
|
||||
content = msg.get("content", "")
|
||||
|
||||
|
||||
# Check for stored output items (from previous Responses API turn)
|
||||
stored_output = msg.get("output")
|
||||
if stored_output and isinstance(stored_output, list):
|
||||
input_items.extend(stored_output)
|
||||
continue
|
||||
|
||||
|
||||
if role == "system":
|
||||
if isinstance(content, str):
|
||||
system_content = content
|
||||
elif isinstance(content, list):
|
||||
system_content = "\n".join(p.get("text", "") for p in content if p.get("type") == "text")
|
||||
system_content = "\n".join(
|
||||
p.get("text", "") for p in content if p.get("type") == "text"
|
||||
)
|
||||
continue
|
||||
|
||||
|
||||
# Convert content format
|
||||
text_type = "output_text" if role == "assistant" else "input_text"
|
||||
|
||||
|
||||
if isinstance(content, str):
|
||||
content_parts = [{"type": text_type, "text": content}]
|
||||
elif isinstance(content, list):
|
||||
content_parts = []
|
||||
for part in content:
|
||||
if part.get("type") == "text":
|
||||
content_parts.append({"type": text_type, "text": part.get("text", "")})
|
||||
content_parts.append(
|
||||
{"type": text_type, "text": part.get("text", "")}
|
||||
)
|
||||
elif part.get("type") == "image_url":
|
||||
url_data = part.get("image_url", {})
|
||||
url = url_data.get("url", "") if isinstance(url_data, dict) else url_data
|
||||
url = (
|
||||
url_data.get("url", "")
|
||||
if isinstance(url_data, dict)
|
||||
else url_data
|
||||
)
|
||||
content_parts.append({"type": "input_image", "image_url": url})
|
||||
else:
|
||||
content_parts = [{"type": text_type, "text": str(content)}]
|
||||
|
||||
input_items.append({
|
||||
"type": "message",
|
||||
"role": role,
|
||||
"content": content_parts
|
||||
})
|
||||
|
||||
|
||||
input_items.append({"type": "message", "role": role, "content": content_parts})
|
||||
|
||||
responses_payload = {**payload, "input": input_items}
|
||||
|
||||
|
||||
if system_content:
|
||||
responses_payload["instructions"] = system_content
|
||||
|
||||
|
||||
if "max_tokens" in responses_payload:
|
||||
responses_payload["max_output_tokens"] = responses_payload.pop("max_tokens")
|
||||
|
||||
|
||||
# Remove Chat Completions-only parameters not supported by the Responses API
|
||||
for unsupported_key in ("stream_options", "logit_bias", "frequency_penalty", "presence_penalty", "stop"):
|
||||
for unsupported_key in (
|
||||
"stream_options",
|
||||
"logit_bias",
|
||||
"frequency_penalty",
|
||||
"presence_penalty",
|
||||
"stop",
|
||||
):
|
||||
responses_payload.pop(unsupported_key, None)
|
||||
|
||||
|
||||
# Convert Chat Completions tools format to Responses API format
|
||||
# Chat Completions: {"type": "function", "function": {"name": ..., "description": ..., "parameters": ...}}
|
||||
# Responses API: {"type": "function", "name": ..., "description": ..., "parameters": ...}
|
||||
@@ -888,9 +897,8 @@ def convert_to_responses_payload(payload: dict) -> dict:
|
||||
# Already in correct format or unknown format, pass through
|
||||
converted_tools.append(tool)
|
||||
responses_payload["tools"] = converted_tools
|
||||
|
||||
return responses_payload
|
||||
|
||||
return responses_payload
|
||||
|
||||
|
||||
def convert_responses_result(response: dict) -> dict:
|
||||
@@ -1036,7 +1044,7 @@ async def generate_chat_completion(
|
||||
headers["api-key"] = key
|
||||
|
||||
headers["api-version"] = api_version
|
||||
|
||||
|
||||
if is_responses:
|
||||
payload = convert_to_responses_payload(payload)
|
||||
request_url = f"{request_url}/responses?api-version={api_version}"
|
||||
|
||||
@@ -107,7 +107,9 @@ async def get_prompt_list(
|
||||
|
||||
filter["user_id"] = user.id
|
||||
|
||||
result = Prompts.search_prompts(user.id, filter=filter, skip=skip, limit=limit, db=db)
|
||||
result = Prompts.search_prompts(
|
||||
user.id, filter=filter, skip=skip, limit=limit, db=db
|
||||
)
|
||||
|
||||
return PromptAccessListResponse(
|
||||
items=[
|
||||
@@ -313,9 +315,7 @@ async def update_prompt_by_id(
|
||||
)
|
||||
|
||||
# Use the ID from the found prompt
|
||||
updated_prompt = Prompts.update_prompt_by_id(
|
||||
prompt.id, form_data, user.id, db=db
|
||||
)
|
||||
updated_prompt = Prompts.update_prompt_by_id(prompt.id, form_data, user.id, db=db)
|
||||
if updated_prompt:
|
||||
return updated_prompt
|
||||
else:
|
||||
@@ -464,9 +464,7 @@ async def update_prompt_access_by_id(
|
||||
detail=ERROR_MESSAGES.ACCESS_PROHIBITED,
|
||||
)
|
||||
|
||||
AccessGrants.set_access_grants(
|
||||
"prompt", prompt_id, form_data.access_grants, db=db
|
||||
)
|
||||
AccessGrants.set_access_grants("prompt", prompt_id, form_data.access_grants, db=db)
|
||||
|
||||
return Prompts.get_prompt_by_id(prompt_id, db=db)
|
||||
|
||||
@@ -522,7 +520,7 @@ async def get_prompt_history(
|
||||
):
|
||||
"""Get version history for a prompt."""
|
||||
PAGE_SIZE = 20
|
||||
|
||||
|
||||
prompt = Prompts.get_prompt_by_id(prompt_id, db=db)
|
||||
|
||||
if not prompt:
|
||||
@@ -554,9 +552,7 @@ async def get_prompt_history(
|
||||
return history
|
||||
|
||||
|
||||
@router.get(
|
||||
"/id/{prompt_id}/history/{history_id}", response_model=PromptHistoryModel
|
||||
)
|
||||
@router.get("/id/{prompt_id}/history/{history_id}", response_model=PromptHistoryModel)
|
||||
async def get_prompt_history_entry(
|
||||
prompt_id: str,
|
||||
history_id: str,
|
||||
@@ -599,9 +595,7 @@ async def get_prompt_history_entry(
|
||||
return history_entry
|
||||
|
||||
|
||||
@router.delete(
|
||||
"/id/{prompt_id}/history/{history_id}", response_model=bool
|
||||
)
|
||||
@router.delete("/id/{prompt_id}/history/{history_id}", response_model=bool)
|
||||
async def delete_prompt_history_entry(
|
||||
prompt_id: str,
|
||||
history_id: str,
|
||||
|
||||
@@ -24,7 +24,6 @@ from open_webui.utils.access_control import has_access, has_permission
|
||||
from open_webui.config import BYPASS_ADMIN_ACCESS_CONTROL
|
||||
from open_webui.constants import ERROR_MESSAGES
|
||||
|
||||
|
||||
log = logging.getLogger(__name__)
|
||||
|
||||
PAGE_ITEM_COUNT = 30
|
||||
@@ -98,9 +97,7 @@ async def get_skill_list(
|
||||
|
||||
filter["user_id"] = user.id
|
||||
|
||||
result = Skills.search_skills(
|
||||
user.id, filter=filter, skip=skip, limit=limit, db=db
|
||||
)
|
||||
result = Skills.search_skills(user.id, filter=filter, skip=skip, limit=limit, db=db)
|
||||
|
||||
return SkillAccessListResponse(
|
||||
items=[
|
||||
@@ -343,9 +340,7 @@ async def update_skill_access_by_id(
|
||||
detail=ERROR_MESSAGES.UNAUTHORIZED,
|
||||
)
|
||||
|
||||
AccessGrants.set_access_grants(
|
||||
"skill", id, form_data.access_grants, db=db
|
||||
)
|
||||
AccessGrants.set_access_grants("skill", id, form_data.access_grants, db=db)
|
||||
|
||||
return Skills.get_skill_by_id(id, db=db)
|
||||
|
||||
|
||||
@@ -36,7 +36,6 @@ from open_webui.config import (
|
||||
DEFAULT_VOICE_MODE_PROMPT_TEMPLATE,
|
||||
)
|
||||
|
||||
|
||||
log = logging.getLogger(__name__)
|
||||
|
||||
router = APIRouter()
|
||||
|
||||
@@ -36,7 +36,6 @@ from open_webui.utils.tools import get_tool_servers
|
||||
from open_webui.config import CACHE_DIR, BYPASS_ADMIN_ACCESS_CONTROL
|
||||
from open_webui.constants import ERROR_MESSAGES
|
||||
|
||||
|
||||
log = logging.getLogger(__name__)
|
||||
|
||||
|
||||
@@ -555,9 +554,7 @@ async def update_tool_access_by_id(
|
||||
detail=ERROR_MESSAGES.UNAUTHORIZED,
|
||||
)
|
||||
|
||||
AccessGrants.set_access_grants(
|
||||
"tool", id, form_data.access_grants, db=db
|
||||
)
|
||||
AccessGrants.set_access_grants("tool", id, form_data.access_grants, db=db)
|
||||
|
||||
return Tools.get_tool_by_id(id, db=db)
|
||||
|
||||
|
||||
@@ -41,7 +41,6 @@ from open_webui.utils.auth import (
|
||||
)
|
||||
from open_webui.utils.access_control import get_permissions, has_permission
|
||||
|
||||
|
||||
log = logging.getLogger(__name__)
|
||||
|
||||
router = APIRouter()
|
||||
|
||||
@@ -15,7 +15,6 @@ from open_webui.utils.pdf_generator import PDFGenerator
|
||||
from open_webui.utils.auth import get_admin_user, get_verified_user
|
||||
from open_webui.utils.code_interpreter import execute_code_jupyter
|
||||
|
||||
|
||||
log = logging.getLogger(__name__)
|
||||
|
||||
router = APIRouter()
|
||||
|
||||
Reference in New Issue
Block a user