diff --git a/backend/open_webui/env.py b/backend/open_webui/env.py index 3fc75c72bb..4289475c04 100644 --- a/backend/open_webui/env.py +++ b/backend/open_webui/env.py @@ -194,6 +194,15 @@ ENABLE_FORWARD_USER_INFO_HEADERS = ( os.environ.get("ENABLE_FORWARD_USER_INFO_HEADERS", "False").lower() == "true" ) +# Header names for user info forwarding (customizable via environment variables) +FORWARD_USER_INFO_HEADER_USER_NAME = os.environ.get("FORWARD_USER_INFO_HEADER_USER_NAME", "X-OpenWebUI-User-Name") +FORWARD_USER_INFO_HEADER_USER_ID = os.environ.get("FORWARD_USER_INFO_HEADER_USER_ID", "X-OpenWebUI-User-Id") +FORWARD_USER_INFO_HEADER_USER_EMAIL = os.environ.get("FORWARD_USER_INFO_HEADER_USER_EMAIL", "X-OpenWebUI-User-Email") +FORWARD_USER_INFO_HEADER_USER_ROLE = os.environ.get("FORWARD_USER_INFO_HEADER_USER_ROLE", "X-OpenWebUI-User-Role") + +# Header name for chat ID forwarding (customizable via environment variable) +FORWARD_SESSION_INFO_HEADER_CHAT_ID = os.environ.get("FORWARD_SESSION_INFO_HEADER_CHAT_ID", "X-OpenWebUI-Chat-Id") + # Experimental feature, may be removed in future ENABLE_STAR_SESSIONS_MIDDLEWARE = ( os.environ.get("ENABLE_STAR_SESSIONS_MIDDLEWARE", "False").lower() == "true" diff --git a/backend/open_webui/functions.py b/backend/open_webui/functions.py index e1604d126a..5381daf0cc 100644 --- a/backend/open_webui/functions.py +++ b/backend/open_webui/functions.py @@ -35,7 +35,6 @@ from open_webui.utils.plugin import ( get_function_module_from_cache, ) from open_webui.utils.tools import get_tools -from open_webui.utils.access_control import has_access from open_webui.env import GLOBAL_LOG_LEVEL diff --git a/backend/open_webui/main.py b/backend/open_webui/main.py index e24807817f..2ff171ebf1 100644 --- a/backend/open_webui/main.py +++ b/backend/open_webui/main.py @@ -518,7 +518,6 @@ from open_webui.utils.middleware import ( process_chat_payload, process_chat_response, ) -from open_webui.utils.access_control import has_access from open_webui.utils.auth import ( get_license_data, @@ -1333,7 +1332,9 @@ class APIKeyRestrictionMiddleware(BaseHTTPMiddleware): token = None if auth_header: - scheme, token = auth_header.split(" ") + parts = auth_header.split(" ", 1) + if len(parts) == 2: + token = parts[1] # Only apply restrictions if an sk- API key is used if token and token.startswith("sk-"): @@ -1806,6 +1807,16 @@ async def chat_completion( except Exception as e: log.debug(f"Error cleaning up: {e}") pass + # Emit chat:active=false when task completes + try: + if metadata.get("chat_id"): + event_emitter = get_event_emitter(metadata, update_db=False) + if event_emitter: + await event_emitter( + {"type": "chat:active", "data": {"active": False}} + ) + except Exception as e: + log.debug(f"Error emitting chat:active: {e}") if ( metadata.get("session_id") @@ -1818,6 +1829,12 @@ async def chat_completion( process_chat(request, form_data, user, metadata, model), id=metadata["chat_id"], ) + # Emit chat:active=true when task starts + event_emitter = get_event_emitter(metadata, update_db=False) + if event_emitter: + await event_emitter( + {"type": "chat:active", "data": {"active": True}} + ) return {"status": True, "task_id": task_id} else: return await process_chat(request, form_data, user, metadata, model) diff --git a/backend/open_webui/migrations/versions/f1e2d3c4b5a6_add_access_grant_table.py b/backend/open_webui/migrations/versions/f1e2d3c4b5a6_add_access_grant_table.py new file mode 100644 index 0000000000..1b76e67c31 --- /dev/null +++ b/backend/open_webui/migrations/versions/f1e2d3c4b5a6_add_access_grant_table.py @@ -0,0 +1,350 @@ +"""Add access_grant table + +Revision ID: f1e2d3c4b5a6 +Revises: 8452d01d26d7 +Create Date: 2026-02-05 10:00:00.000000 + +Migrates from JSON access_control columns to normalized access_grant table. +Access control semantics: +- NULL: Public access (all users can read) -> insert user:* for read +- {}: Private/owner-only (no grants) -> insert nothing +- {read: {...}, write: {...}}: Custom permissions -> insert specific grants +""" + +from typing import Sequence, Union +import time +import uuid + +from alembic import op +import sqlalchemy as sa + +from open_webui.migrations.util import get_existing_tables + +revision: str = "f1e2d3c4b5a6" +down_revision: Union[str, None] = "8452d01d26d7" +branch_labels: Union[str, Sequence[str], None] = None +depends_on: Union[str, Sequence[str], None] = None + + +def upgrade() -> None: + existing_tables = set(get_existing_tables()) + + # Create access_grant table + if "access_grant" not in existing_tables: + op.create_table( + "access_grant", + sa.Column("id", sa.Text(), nullable=False, primary_key=True), + sa.Column("resource_type", sa.Text(), nullable=False), + sa.Column("resource_id", sa.Text(), nullable=False), + sa.Column("principal_type", sa.Text(), nullable=False), + sa.Column("principal_id", sa.Text(), nullable=False), + sa.Column("permission", sa.Text(), nullable=False), + sa.Column("created_at", sa.BigInteger(), nullable=False), + sa.UniqueConstraint( + "resource_type", + "resource_id", + "principal_type", + "principal_id", + "permission", + name="uq_access_grant_grant", + ), + ) + op.create_index( + "idx_access_grant_resource", + "access_grant", + ["resource_type", "resource_id"], + ) + op.create_index( + "idx_access_grant_principal", + "access_grant", + ["principal_type", "principal_id"], + ) + + # Backfill existing access_control JSON data + conn = op.get_bind() + + # Tables with access_control JSON columns: (table_name, resource_type) + resource_tables = [ + ("knowledge", "knowledge"), + ("prompt", "prompt"), + ("tool", "tool"), + ("model", "model"), + ("note", "note"), + ("channel", "channel"), + ("file", "file"), + ] + + now = int(time.time()) + inserted = set() + + for table_name, resource_type in resource_tables: + if table_name not in existing_tables: + continue + + # Query all rows + try: + result = conn.execute( + sa.text(f'SELECT id, access_control FROM "{table_name}"') + ) + rows = result.fetchall() + except Exception: + continue + + for row in rows: + resource_id = row[0] + access_control_json = row[1] + + # Handle NULL or JSON "null" = public access (user:* for read) + # Could be Python None (SQL NULL) or string "null" (JSON null) + # EXCEPTION: files with NULL are PRIVATE (owner-only), not public + is_null = ( + access_control_json is None or + access_control_json == "null" or + (isinstance(access_control_json, str) and access_control_json.strip().lower() == "null") + ) + if is_null: + # Files: NULL = private (no entry needed, owner has implicit access) + # Other resources: NULL = public (insert user:* for read) + if resource_type == "file": + continue # Private - no entry needed + + key = (resource_type, resource_id, "user", "*", "read") + if key not in inserted: + try: + conn.execute( + sa.text( + """ + INSERT INTO access_grant (id, resource_type, resource_id, principal_type, principal_id, permission, created_at) + VALUES (:id, :resource_type, :resource_id, :principal_type, :principal_id, :permission, :created_at) + """ + ), + { + "id": str(uuid.uuid4()), + "resource_type": resource_type, + "resource_id": resource_id, + "principal_type": "user", + "principal_id": "*", + "permission": "read", + "created_at": now, + }, + ) + inserted.add(key) + except Exception: + pass + continue + + # Handle JSON parsing + if isinstance(access_control_json, str): + import json + + try: + access_control_json = json.loads(access_control_json) + except Exception: + continue + + # Handle {} = private/owner-only - NO entries needed + # Owner access is implicit, no grants to store + if not access_control_json or not isinstance(access_control_json, dict): + continue + + # Check if it's effectively empty (no read/write keys with content) + read_data = access_control_json.get("read", {}) + write_data = access_control_json.get("write", {}) + + has_read_grants = read_data.get("group_ids", []) or read_data.get( + "user_ids", [] + ) + has_write_grants = write_data.get("group_ids", []) or write_data.get( + "user_ids", [] + ) + + if not has_read_grants and not has_write_grants: + # Empty permissions = private, no grants needed + continue + + # Extract permissions and insert into access_grant table + for permission in ["read", "write"]: + perm_data = access_control_json.get(permission, {}) + if not perm_data: + continue + + for group_id in perm_data.get("group_ids", []): + key = (resource_type, resource_id, "group", group_id, permission) + if key in inserted: + continue + try: + conn.execute( + sa.text( + """ + INSERT INTO access_grant (id, resource_type, resource_id, principal_type, principal_id, permission, created_at) + VALUES (:id, :resource_type, :resource_id, :principal_type, :principal_id, :permission, :created_at) + """ + ), + { + "id": str(uuid.uuid4()), + "resource_type": resource_type, + "resource_id": resource_id, + "principal_type": "group", + "principal_id": group_id, + "permission": permission, + "created_at": now, + }, + ) + inserted.add(key) + except Exception: + pass + + for user_id in perm_data.get("user_ids", []): + key = (resource_type, resource_id, "user", user_id, permission) + if key in inserted: + continue + try: + conn.execute( + sa.text( + """ + INSERT INTO access_grant (id, resource_type, resource_id, principal_type, principal_id, permission, created_at) + VALUES (:id, :resource_type, :resource_id, :principal_type, :principal_id, :permission, :created_at) + """ + ), + { + "id": str(uuid.uuid4()), + "resource_type": resource_type, + "resource_id": resource_id, + "principal_type": "user", + "principal_id": user_id, + "permission": permission, + "created_at": now, + }, + ) + inserted.add(key) + except Exception: + pass + + # Drop access_control columns from resource tables + for table_name, _ in resource_tables: + if table_name not in existing_tables: + continue + try: + with op.batch_alter_table(table_name) as batch: + batch.drop_column("access_control") + except Exception: + pass + + +def downgrade() -> None: + import json + + conn = op.get_bind() + + # Resource tables mapping: (table_name, resource_type) + resource_tables = [ + ("knowledge", "knowledge"), + ("prompt", "prompt"), + ("tool", "tool"), + ("model", "model"), + ("note", "note"), + ("channel", "channel"), + ("file", "file"), + ] + + # Step 1: Re-add access_control columns to resource tables + for table_name, _ in resource_tables: + try: + with op.batch_alter_table(table_name) as batch: + batch.add_column(sa.Column("access_control", sa.JSON(), nullable=True)) + except Exception: + pass + + # Step 2: Query access_grant table and reconstruct JSON for each resource + for table_name, resource_type in resource_tables: + try: + # Get all grants for this resource type + result = conn.execute( + sa.text(""" + SELECT resource_id, principal_type, principal_id, permission + FROM access_grant + WHERE resource_type = :resource_type + """), + {"resource_type": resource_type} + ) + rows = result.fetchall() + except Exception: + continue + + # Group by resource_id and reconstruct JSON structure + resource_grants = {} + for row in rows: + resource_id = row[0] + principal_type = row[1] + principal_id = row[2] + permission = row[3] + + if resource_id not in resource_grants: + resource_grants[resource_id] = { + "is_public": False, + "read": {"group_ids": [], "user_ids": []}, + "write": {"group_ids": [], "user_ids": []}, + } + + # Handle public access (user:* for read) + if principal_type == "user" and principal_id == "*" and permission == "read": + resource_grants[resource_id]["is_public"] = True + continue + + # Add to appropriate list + if permission in ["read", "write"]: + if principal_type == "group": + if principal_id not in resource_grants[resource_id][permission]["group_ids"]: + resource_grants[resource_id][permission]["group_ids"].append(principal_id) + elif principal_type == "user": + if principal_id not in resource_grants[resource_id][permission]["user_ids"]: + resource_grants[resource_id][permission]["user_ids"].append(principal_id) + + # Step 3: Update each resource with reconstructed JSON + for resource_id, grants in resource_grants.items(): + if grants["is_public"]: + # Public = NULL + access_control_value = None + elif (not grants["read"]["group_ids"] and not grants["read"]["user_ids"] and + not grants["write"]["group_ids"] and not grants["write"]["user_ids"]): + # No grants = should not happen (would mean no entries), default to {} + access_control_value = json.dumps({}) + else: + # Custom permissions + access_control_value = json.dumps({ + "read": grants["read"], + "write": grants["write"], + }) + + try: + conn.execute( + sa.text(f'UPDATE "{table_name}" SET access_control = :access_control WHERE id = :id'), + {"access_control": access_control_value, "id": resource_id} + ) + except Exception: + pass + + # Step 4: Set all resources WITHOUT entries to private + # For files: NULL means private (owner-only), so leave as NULL + # For other resources: {} means private, so update to {} + if resource_type != "file": + try: + conn.execute( + sa.text(f''' + UPDATE "{table_name}" + SET access_control = :private_value + WHERE id NOT IN ( + SELECT DISTINCT resource_id FROM access_grant WHERE resource_type = :resource_type + ) + AND access_control IS NULL + '''), + {"private_value": json.dumps({}), "resource_type": resource_type} + ) + except Exception: + pass + # For files, NULL stays NULL - no action needed + + # Step 5: Drop the access_grant table + op.drop_index("idx_access_grant_principal", table_name="access_grant") + op.drop_index("idx_access_grant_resource", table_name="access_grant") + op.drop_table("access_grant") diff --git a/backend/open_webui/models/access_grants.py b/backend/open_webui/models/access_grants.py new file mode 100644 index 0000000000..aac475f3e1 --- /dev/null +++ b/backend/open_webui/models/access_grants.py @@ -0,0 +1,776 @@ +import logging +import time +import uuid +from typing import Optional + +from sqlalchemy.orm import Session +from open_webui.internal.db import Base, get_db_context + +from pydantic import BaseModel, ConfigDict +from sqlalchemy import BigInteger, Column, Text, UniqueConstraint, or_, and_ +from sqlalchemy.dialects.postgresql import JSONB + +log = logging.getLogger(__name__) + + +#################### +# AccessGrant DB Schema +#################### + + +class AccessGrant(Base): + __tablename__ = "access_grant" + + id = Column(Text, primary_key=True) + resource_type = Column(Text, nullable=False) # "knowledge", "model", "prompt", "tool", "note", "channel", "file" + resource_id = Column(Text, nullable=False) + principal_type = Column(Text, nullable=False) # "user" or "group" + principal_id = Column(Text, nullable=False) # user_id, group_id, or "*" (wildcard for public) + permission = Column(Text, nullable=False) # "read" or "write" + created_at = Column(BigInteger, nullable=False) + + __table_args__ = ( + UniqueConstraint( + "resource_type", + "resource_id", + "principal_type", + "principal_id", + "permission", + name="uq_access_grant_grant", + ), + ) + + +class AccessGrantModel(BaseModel): + model_config = ConfigDict(from_attributes=True) + + id: str + resource_type: str + resource_id: str + principal_type: str + principal_id: str + permission: str + created_at: int + + +class AccessGrantResponse(BaseModel): + """Slim grant model for API responses — resource context is implicit from the parent.""" + + id: str + principal_type: str + principal_id: str + permission: str + + @classmethod + def from_grant(cls, grant: "AccessGrantModel") -> "AccessGrantResponse": + return cls( + id=grant.id, + principal_type=grant.principal_type, + principal_id=grant.principal_id, + permission=grant.permission, + ) + + +#################### +# Conversion utilities +#################### + + +def access_control_to_grants( + resource_type: str, + resource_id: str, + access_control: Optional[dict], +) -> list[dict]: + """ + Convert an old-style access_control JSON dict to a flat list of grant dicts. + + Semantics: + - None → public read (user:* read) — except files which are private + - {} → private/owner-only (no grants) + - {read: {group_ids, user_ids}, write: {group_ids, user_ids}} → specific grants + + Returns a list of dicts with keys: resource_type, resource_id, principal_type, principal_id, permission + """ + grants = [] + + if access_control is None: + # NULL → public read (user:* for read) + # Exception: files with NULL are private (owner-only), no grants needed + if resource_type != "file": + grants.append( + { + "resource_type": resource_type, + "resource_id": resource_id, + "principal_type": "user", + "principal_id": "*", + "permission": "read", + } + ) + return grants + + # {} → private/owner-only, no grants + if not access_control: + return grants + + # Parse structured permissions + for permission in ["read", "write"]: + perm_data = access_control.get(permission, {}) + if not perm_data: + continue + + for group_id in perm_data.get("group_ids", []): + grants.append( + { + "resource_type": resource_type, + "resource_id": resource_id, + "principal_type": "group", + "principal_id": group_id, + "permission": permission, + } + ) + + for user_id in perm_data.get("user_ids", []): + grants.append( + { + "resource_type": resource_type, + "resource_id": resource_id, + "principal_type": "user", + "principal_id": user_id, + "permission": permission, + } + ) + + return grants + + +def normalize_access_grants(access_grants: Optional[list]) -> list[dict]: + """ + Normalize direct access_grants payloads from API forms. + + Keeps only valid grants and removes duplicates by + (principal_type, principal_id, permission). + """ + if not access_grants: + return [] + + deduped = {} + for grant in access_grants: + if isinstance(grant, BaseModel): + grant = grant.model_dump() + if not isinstance(grant, dict): + continue + + principal_type = grant.get("principal_type") + principal_id = grant.get("principal_id") + permission = grant.get("permission") + + if principal_type not in ("user", "group"): + continue + if permission not in ("read", "write"): + continue + if not isinstance(principal_id, str) or not principal_id: + continue + + key = (principal_type, principal_id, permission) + deduped[key] = { + "id": grant.get("id") + if isinstance(grant.get("id"), str) and grant.get("id") + else str(uuid.uuid4()), + "principal_type": principal_type, + "principal_id": principal_id, + "permission": permission, + } + + return list(deduped.values()) + + +def has_public_read_access_grant(access_grants: Optional[list]) -> bool: + """ + Returns True when a direct grant list includes wildcard public-read. + """ + for grant in normalize_access_grants(access_grants): + if ( + grant["principal_type"] == "user" + and grant["principal_id"] == "*" + and grant["permission"] == "read" + ): + return True + return False + + +def grants_to_access_control(grants: list) -> Optional[dict]: + """ + Convert a list of grant objects (AccessGrantModel or AccessGrantResponse) + back to the old-style access_control JSON dict for backward compatibility. + + Semantics: + - [] (empty) → {} (private/owner-only) + - Contains user:*:read → None (public), but write grants are preserved + - Otherwise → {read: {group_ids, user_ids}, write: {group_ids, user_ids}} + + Note: "public" (user:*:read) still allows additional write permissions + to coexist. When the wildcard read is present the function returns None + for the legacy dict, so callers that need write info should inspect the + grants list directly. + """ + if not grants: + return {} # No grants = private/owner-only + + result = { + "read": {"group_ids": [], "user_ids": []}, + "write": {"group_ids": [], "user_ids": []}, + } + + is_public = False + for grant in grants: + if ( + grant.principal_type == "user" + and grant.principal_id == "*" + and grant.permission == "read" + ): + is_public = True + continue # Don't add wildcard to user_ids list + + if grant.permission not in ("read", "write"): + continue + + if grant.principal_type == "group": + if grant.principal_id not in result[grant.permission]["group_ids"]: + result[grant.permission]["group_ids"].append(grant.principal_id) + elif grant.principal_type == "user": + if grant.principal_id not in result[grant.permission]["user_ids"]: + result[grant.permission]["user_ids"].append(grant.principal_id) + + if is_public: + return None # Public read access + + return result + + +#################### +# Table Operations +#################### + + +class AccessGrantsTable: + def grant_access( + self, + resource_type: str, + resource_id: str, + principal_type: str, + principal_id: str, + permission: str, + db: Optional[Session] = None, + ) -> Optional[AccessGrantModel]: + """Add a single access grant. Idempotent (ignores duplicates).""" + with get_db_context(db) as db: + # Check for existing grant + existing = ( + db.query(AccessGrant) + .filter_by( + resource_type=resource_type, + resource_id=resource_id, + principal_type=principal_type, + principal_id=principal_id, + permission=permission, + ) + .first() + ) + if existing: + return AccessGrantModel.model_validate(existing) + + grant = AccessGrant( + id=str(uuid.uuid4()), + resource_type=resource_type, + resource_id=resource_id, + principal_type=principal_type, + principal_id=principal_id, + permission=permission, + created_at=int(time.time()), + ) + db.add(grant) + db.commit() + db.refresh(grant) + return AccessGrantModel.model_validate(grant) + + def revoke_access( + self, + resource_type: str, + resource_id: str, + principal_type: str, + principal_id: str, + permission: str, + db: Optional[Session] = None, + ) -> bool: + """Remove a single access grant.""" + with get_db_context(db) as db: + deleted = ( + db.query(AccessGrant) + .filter_by( + resource_type=resource_type, + resource_id=resource_id, + principal_type=principal_type, + principal_id=principal_id, + permission=permission, + ) + .delete() + ) + db.commit() + return deleted > 0 + + def revoke_all_access( + self, + resource_type: str, + resource_id: str, + db: Optional[Session] = None, + ) -> int: + """Remove all access grants for a resource.""" + with get_db_context(db) as db: + deleted = ( + db.query(AccessGrant) + .filter_by( + resource_type=resource_type, + resource_id=resource_id, + ) + .delete() + ) + db.commit() + return deleted + + def set_access_control( + self, + resource_type: str, + resource_id: str, + access_control: Optional[dict], + db: Optional[Session] = None, + ) -> list[AccessGrantModel]: + """ + Replace all grants for a resource from an access_control JSON dict. + This is the primary bridge for backward compat with the frontend. + """ + with get_db_context(db) as db: + # Delete all existing grants for this resource + db.query(AccessGrant).filter_by( + resource_type=resource_type, + resource_id=resource_id, + ).delete() + + # Convert JSON to grant dicts + grant_dicts = access_control_to_grants( + resource_type, resource_id, access_control + ) + + # Insert new grants + results = [] + for grant_dict in grant_dicts: + grant = AccessGrant( + id=str(uuid.uuid4()), + **grant_dict, + created_at=int(time.time()), + ) + db.add(grant) + results.append(grant) + + db.commit() + + return [AccessGrantModel.model_validate(g) for g in results] + + def set_access_grants( + self, + resource_type: str, + resource_id: str, + access_grants: Optional[list], + db: Optional[Session] = None, + ) -> list[AccessGrantModel]: + """ + Replace all grants for a resource from a direct access_grants list. + """ + with get_db_context(db) as db: + db.query(AccessGrant).filter_by( + resource_type=resource_type, + resource_id=resource_id, + ).delete() + + normalized_grants = normalize_access_grants(access_grants) + + results = [] + for grant_dict in normalized_grants: + grant = AccessGrant( + id=grant_dict["id"], + resource_type=resource_type, + resource_id=resource_id, + principal_type=grant_dict["principal_type"], + principal_id=grant_dict["principal_id"], + permission=grant_dict["permission"], + created_at=int(time.time()), + ) + db.add(grant) + results.append(grant) + + db.commit() + return [AccessGrantModel.model_validate(g) for g in results] + + def get_access_control( + self, + resource_type: str, + resource_id: str, + db: Optional[Session] = None, + ) -> Optional[dict]: + """ + Reconstruct the old-style access_control JSON dict from grants. + For backward compat with the frontend. + """ + with get_db_context(db) as db: + grants = ( + db.query(AccessGrant) + .filter_by( + resource_type=resource_type, + resource_id=resource_id, + ) + .all() + ) + grant_models = [AccessGrantModel.model_validate(g) for g in grants] + return grants_to_access_control(grant_models) + + def get_grants_by_resource( + self, + resource_type: str, + resource_id: str, + db: Optional[Session] = None, + ) -> list[AccessGrantModel]: + """Get all grants for a specific resource.""" + with get_db_context(db) as db: + grants = ( + db.query(AccessGrant) + .filter_by( + resource_type=resource_type, + resource_id=resource_id, + ) + .all() + ) + return [AccessGrantModel.model_validate(g) for g in grants] + + def has_access( + self, + user_id: str, + resource_type: str, + resource_id: str, + permission: str = "read", + user_group_ids: Optional[set[str]] = None, + db: Optional[Session] = None, + ) -> bool: + """ + Check if a user has the specified permission on a resource. + + Access is granted if any of the following is true: + - There's a grant for user:* (public) with the requested permission + - There's a grant for the specific user with the requested permission + - There's a grant for any of the user's groups with the requested permission + """ + with get_db_context(db) as db: + # Build conditions for matching grants + conditions = [ + # Public access + and_( + AccessGrant.principal_type == "user", + AccessGrant.principal_id == "*", + ), + # Direct user access + and_( + AccessGrant.principal_type == "user", + AccessGrant.principal_id == user_id, + ), + ] + + # Group access + if user_group_ids is None: + from open_webui.models.groups import Groups + + user_groups = Groups.get_groups_by_member_id(user_id, db=db) + user_group_ids = {group.id for group in user_groups} + + if user_group_ids: + conditions.append( + and_( + AccessGrant.principal_type == "group", + AccessGrant.principal_id.in_(user_group_ids), + ) + ) + + exists = ( + db.query(AccessGrant) + .filter( + AccessGrant.resource_type == resource_type, + AccessGrant.resource_id == resource_id, + AccessGrant.permission == permission, + or_(*conditions), + ) + .first() + ) + return exists is not None + + def get_users_with_access( + self, + resource_type: str, + resource_id: str, + permission: str = "read", + db: Optional[Session] = None, + ) -> list: + """ + Get all users who have the specified permission on a resource. + Returns a list of UserModel instances. + """ + from open_webui.models.users import Users, UserModel + from open_webui.models.groups import Groups + + with get_db_context(db) as db: + grants = ( + db.query(AccessGrant) + .filter_by( + resource_type=resource_type, + resource_id=resource_id, + permission=permission, + ) + .all() + ) + + # Check for public access + for grant in grants: + if grant.principal_type == "user" and grant.principal_id == "*": + result = Users.get_users(filter={"roles": ["!pending"]}, db=db) + return result.get("users", []) + + user_ids_with_access = set() + + for grant in grants: + if grant.principal_type == "user": + user_ids_with_access.add(grant.principal_id) + elif grant.principal_type == "group": + group_user_ids = Groups.get_group_user_ids_by_id( + grant.principal_id, db=db + ) + if group_user_ids: + user_ids_with_access.update(group_user_ids) + + if not user_ids_with_access: + return [] + + return Users.get_users_by_user_ids(list(user_ids_with_access), db=db) + + def has_permission_filter( + self, + db, + query, + DocumentModel, + filter: dict, + resource_type: str, + permission: str = "read", + ): + """ + Apply access control filtering to a SQLAlchemy query by JOINing with access_grant. + + This replaces the old JSON-column-based filtering with a proper relational JOIN. + """ + group_ids = filter.get("group_ids", []) + user_id = filter.get("user_id") + + if permission == "read_only": + return self._has_read_only_permission_filter( + db, query, DocumentModel, filter, resource_type + ) + + # Build principal conditions + principal_conditions = [] + + if group_ids or user_id: + # Public access: user:* read + principal_conditions.append( + and_( + AccessGrant.principal_type == "user", + AccessGrant.principal_id == "*", + ) + ) + + if user_id: + # Owner always has access + principal_conditions.append(DocumentModel.user_id == user_id) + + # Direct user grant + principal_conditions.append( + and_( + AccessGrant.principal_type == "user", + AccessGrant.principal_id == user_id, + ) + ) + + if group_ids: + # Group grants + principal_conditions.append( + and_( + AccessGrant.principal_type == "group", + AccessGrant.principal_id.in_(group_ids), + ) + ) + + if not principal_conditions: + return query + + # LEFT JOIN access_grant and filter + # We use a subquery approach to avoid duplicates from multiple matching grants + from sqlalchemy import exists as sa_exists, select + + grant_exists = ( + select(AccessGrant.id) + .where( + AccessGrant.resource_type == resource_type, + AccessGrant.resource_id == DocumentModel.id, + AccessGrant.permission == permission, + or_( + and_( + AccessGrant.principal_type == "user", + AccessGrant.principal_id == "*", + ), + *( + [ + and_( + AccessGrant.principal_type == "user", + AccessGrant.principal_id == user_id, + ) + ] + if user_id + else [] + ), + *( + [ + and_( + AccessGrant.principal_type == "group", + AccessGrant.principal_id.in_(group_ids), + ) + ] + if group_ids + else [] + ), + ), + ) + .correlate(DocumentModel) + .exists() + ) + + # Owner OR has a matching grant + owner_or_grant = [grant_exists] + if user_id: + owner_or_grant.append(DocumentModel.user_id == user_id) + + query = query.filter(or_(*owner_or_grant)) + return query + + def _has_read_only_permission_filter( + self, + db, + query, + DocumentModel, + filter: dict, + resource_type: str, + ): + """ + Filter for items where user has read BUT NOT write access. + Public items are NOT considered read_only. + """ + group_ids = filter.get("group_ids", []) + user_id = filter.get("user_id") + + from sqlalchemy import exists as sa_exists, select + + # Has read grant (not public) + read_grant_exists = ( + select(AccessGrant.id) + .where( + AccessGrant.resource_type == resource_type, + AccessGrant.resource_id == DocumentModel.id, + AccessGrant.permission == "read", + or_( + *( + [ + and_( + AccessGrant.principal_type == "user", + AccessGrant.principal_id == user_id, + ) + ] + if user_id + else [] + ), + *( + [ + and_( + AccessGrant.principal_type == "group", + AccessGrant.principal_id.in_(group_ids), + ) + ] + if group_ids + else [] + ), + ), + ) + .correlate(DocumentModel) + .exists() + ) + + # Does NOT have write grant + write_grant_exists = ( + select(AccessGrant.id) + .where( + AccessGrant.resource_type == resource_type, + AccessGrant.resource_id == DocumentModel.id, + AccessGrant.permission == "write", + or_( + *( + [ + and_( + AccessGrant.principal_type == "user", + AccessGrant.principal_id == user_id, + ) + ] + if user_id + else [] + ), + *( + [ + and_( + AccessGrant.principal_type == "group", + AccessGrant.principal_id.in_(group_ids), + ) + ] + if group_ids + else [] + ), + ), + ) + .correlate(DocumentModel) + .exists() + ) + + # Is NOT public + public_grant_exists = ( + select(AccessGrant.id) + .where( + AccessGrant.resource_type == resource_type, + AccessGrant.resource_id == DocumentModel.id, + AccessGrant.permission == "read", + AccessGrant.principal_type == "user", + AccessGrant.principal_id == "*", + ) + .correlate(DocumentModel) + .exists() + ) + + conditions = [read_grant_exists, ~write_grant_exists, ~public_grant_exists] + + # Not owner + if user_id: + conditions.append(DocumentModel.user_id != user_id) + + query = query.filter(and_(*conditions)) + return query + + +AccessGrants = AccessGrantsTable() diff --git a/backend/open_webui/models/channels.py b/backend/open_webui/models/channels.py index 8e70918e1a..3ff6fb7554 100644 --- a/backend/open_webui/models/channels.py +++ b/backend/open_webui/models/channels.py @@ -7,8 +7,12 @@ from typing import Optional from sqlalchemy.orm import Session from open_webui.internal.db import Base, JSONField, get_db, get_db_context from open_webui.models.groups import Groups +from open_webui.models.access_grants import ( + AccessGrantModel, + AccessGrants, +) -from pydantic import BaseModel, ConfigDict +from pydantic import BaseModel, ConfigDict, Field from sqlalchemy.dialects.postgresql import JSONB @@ -47,7 +51,6 @@ class Channel(Base): data = Column(JSON, nullable=True) meta = Column(JSON, nullable=True) - access_control = Column(JSON, nullable=True) created_at = Column(BigInteger) @@ -76,7 +79,7 @@ class ChannelModel(BaseModel): data: Optional[dict] = None meta: Optional[dict] = None - access_control: Optional[dict] = None + access_grants: list[AccessGrantModel] = Field(default_factory=list) created_at: int # timestamp in epoch (time_ns) @@ -237,7 +240,7 @@ class ChannelForm(BaseModel): is_private: Optional[bool] = None data: Optional[dict] = None meta: Optional[dict] = None - access_control: Optional[dict] = None + access_grants: Optional[list[dict]] = None group_ids: Optional[list[str]] = None user_ids: Optional[list[str]] = None @@ -252,6 +255,18 @@ class ChannelWebhookForm(BaseModel): class ChannelTable: + def _get_access_grants( + self, channel_id: str, db: Optional[Session] = None + ) -> list[AccessGrantModel]: + return AccessGrants.get_grants_by_resource("channel", channel_id, db=db) + + def _to_channel_model( + self, channel: Channel, db: Optional[Session] = None + ) -> ChannelModel: + channel_data = ChannelModel.model_validate(channel).model_dump(exclude={"access_grants"}) + access_grants = self._get_access_grants(channel_data["id"], db=db) + channel_data["access_grants"] = access_grants + return ChannelModel.model_validate(channel_data) def _collect_unique_user_ids( self, @@ -316,16 +331,17 @@ class ChannelTable: with get_db_context(db) as db: channel = ChannelModel( **{ - **form_data.model_dump(), + **form_data.model_dump(exclude={"access_grants"}), "type": form_data.type if form_data.type else None, "name": form_data.name.lower(), "id": str(uuid.uuid4()), "user_id": user_id, "created_at": int(time.time_ns()), "updated_at": int(time.time_ns()), + "access_grants": [], } ) - new_channel = Channel(**channel.model_dump()) + new_channel = Channel(**channel.model_dump(exclude={"access_grants"})) if form_data.type in ["group", "dm"]: users = self._collect_unique_user_ids( @@ -342,54 +358,25 @@ class ChannelTable: db.add_all(memberships) db.add(new_channel) db.commit() - return channel + AccessGrants.set_access_grants( + "channel", new_channel.id, form_data.access_grants, db=db + ) + return self._to_channel_model(new_channel, db=db) def get_channels(self, db: Optional[Session] = None) -> list[ChannelModel]: with get_db_context(db) as db: channels = db.query(Channel).all() - return [ChannelModel.model_validate(channel) for channel in channels] + return [self._to_channel_model(channel, db=db) for channel in channels] def _has_permission(self, db, query, filter: dict, permission: str = "read"): - group_ids = filter.get("group_ids", []) - user_id = filter.get("user_id") - - dialect_name = db.bind.dialect.name - - # Public access - conditions = [] - if group_ids or user_id: - conditions.extend( - [ - Channel.access_control.is_(None), - cast(Channel.access_control, String) == "null", - ] - ) - - # User-level permission - if user_id: - conditions.append(Channel.user_id == user_id) - - # Group-level permission - if group_ids: - group_conditions = [] - for gid in group_ids: - if dialect_name == "sqlite": - group_conditions.append( - Channel.access_control[permission]["group_ids"].contains([gid]) - ) - elif dialect_name == "postgresql": - group_conditions.append( - cast( - Channel.access_control[permission]["group_ids"], - JSONB, - ).contains([gid]) - ) - conditions.append(or_(*group_conditions)) - - if conditions: - query = query.filter(or_(*conditions)) - - return query + return AccessGrants.has_permission_filter( + db=db, + query=query, + DocumentModel=Channel, + filter=filter, + resource_type="channel", + permission=permission, + ) def get_channels_by_user_id( self, user_id: str, db: Optional[Session] = None @@ -428,7 +415,7 @@ class ChannelTable: standard_channels = query.all() all_channels = membership_channels + standard_channels - return [ChannelModel.model_validate(c) for c in all_channels] + return [self._to_channel_model(c, db=db) for c in all_channels] def get_dm_channel_by_user_ids( self, user_ids: list[str], db: Optional[Session] = None @@ -463,7 +450,7 @@ class ChannelTable: .first() ) - return ChannelModel.model_validate(channel) if channel else None + return self._to_channel_model(channel, db=db) if channel else None def add_members_to_channel( self, @@ -722,7 +709,7 @@ class ChannelTable: try: with get_db_context(db) as db: channel = db.query(Channel).filter(Channel.id == id).first() - return ChannelModel.model_validate(channel) if channel else None + return self._to_channel_model(channel, db=db) if channel else None except Exception: return None @@ -735,7 +722,7 @@ class ChannelTable: ) channel_ids = [cf.channel_id for cf in channel_files] channels = db.query(Channel).filter(Channel.id.in_(channel_ids)).all() - return [ChannelModel.model_validate(channel) for channel in channels] + return [self._to_channel_model(channel, db=db) for channel in channels] def get_channels_by_file_id_and_user_id( self, file_id: str, user_id: str, db: Optional[Session] = None @@ -783,7 +770,9 @@ class ChannelTable: .first() ) if membership: - allowed_channels.append(ChannelModel.model_validate(channel)) + allowed_channels.append( + self._to_channel_model(channel, db=db) + ) continue # --- Case B: standard channel => rely on ACL permissions --- @@ -798,7 +787,7 @@ class ChannelTable: allowed = query.first() if allowed: - allowed_channels.append(ChannelModel.model_validate(allowed)) + allowed_channels.append(self._to_channel_model(allowed, db=db)) return allowed_channels @@ -832,7 +821,7 @@ class ChannelTable: .first() ) if membership: - return ChannelModel.model_validate(channel) + return self._to_channel_model(channel, db=db) else: return None @@ -854,7 +843,7 @@ class ChannelTable: channel_allowed = query.first() return ( - ChannelModel.model_validate(channel_allowed) + self._to_channel_model(channel_allowed, db=db) if channel_allowed else None ) @@ -874,11 +863,14 @@ class ChannelTable: channel.data = form_data.data channel.meta = form_data.meta - channel.access_control = form_data.access_control + if form_data.access_grants is not None: + AccessGrants.set_access_grants( + "channel", id, form_data.access_grants, db=db + ) channel.updated_at = int(time.time_ns()) db.commit() - return ChannelModel.model_validate(channel) if channel else None + return self._to_channel_model(channel, db=db) if channel else None def add_file_to_channel_by_id( self, channel_id: str, file_id: str, user_id: str, db: Optional[Session] = None @@ -947,6 +939,7 @@ class ChannelTable: def delete_channel_by_id(self, id: str, db: Optional[Session] = None) -> bool: with get_db_context(db) as db: + AccessGrants.revoke_all_access("channel", id, db=db) db.query(Channel).filter(Channel.id == id).delete() db.commit() return True diff --git a/backend/open_webui/models/files.py b/backend/open_webui/models/files.py index c24b242bd8..67f2891605 100644 --- a/backend/open_webui/models/files.py +++ b/backend/open_webui/models/files.py @@ -26,8 +26,6 @@ class File(Base): data = Column(JSON, nullable=True) meta = Column(JSON, nullable=True) - access_control = Column(JSON, nullable=True) - created_at = Column(BigInteger) updated_at = Column(BigInteger) @@ -45,8 +43,6 @@ class FileModel(BaseModel): data: Optional[dict] = None meta: Optional[dict] = None - access_control: Optional[dict] = None - created_at: Optional[int] # timestamp in epoch updated_at: Optional[int] # timestamp in epoch @@ -113,7 +109,6 @@ class FileForm(BaseModel): path: str data: dict = {} meta: dict = {} - access_control: Optional[dict] = None class FileUpdateForm(BaseModel): diff --git a/backend/open_webui/models/groups.py b/backend/open_webui/models/groups.py index eab817534b..0859c053aa 100644 --- a/backend/open_webui/models/groups.py +++ b/backend/open_webui/models/groups.py @@ -22,6 +22,7 @@ from sqlalchemy import ( ForeignKey, cast, or_, + select, ) @@ -99,6 +100,16 @@ class GroupResponse(GroupModel): member_count: Optional[int] = None +class GroupInfoResponse(BaseModel): + id: str + user_id: str + name: str + description: str + member_count: Optional[int] = None + created_at: int + updated_at: int + + class GroupForm(BaseModel): name: str description: str @@ -171,22 +182,22 @@ class GroupTable: if share_value: # Groups open to anyone: data is null, config.share is null, or share is true # Use case-insensitive string comparison to handle variations like "True", "TRUE" + # Handle potential JSON boolean to string casting issues by checking for both string 'true' and boolean equivalence if possible, anyone_can_share = or_( Group.data.is_(None), json_share_str.is_(None), json_share_lower == "true", + json_share_lower == "1", # Handle SQLite boolean true ) if member_id: # Also include member-only groups where user is a member - member_groups_subq = ( - db.query(GroupMember.group_id) - .filter(GroupMember.user_id == member_id) - .subquery() + member_groups_select = select(GroupMember.group_id).where( + GroupMember.user_id == member_id ) members_only_and_is_member = and_( json_share_lower == "members", - Group.id.in_(member_groups_subq), + Group.id.in_(member_groups_select), ) query = query.filter( or_(anyone_can_share, members_only_and_is_member) @@ -305,14 +316,14 @@ class GroupTable: def get_group_user_ids_by_id( self, id: str, db: Optional[Session] = None - ) -> Optional[list[str]]: + ) -> list[str]: with get_db_context(db) as db: members = ( db.query(GroupMember.user_id).filter(GroupMember.group_id == id).all() ) if not members: - return None + return [] return [m[0] for m in members] diff --git a/backend/open_webui/models/knowledge.py b/backend/open_webui/models/knowledge.py index 81aa4099d9..817cab5caf 100644 --- a/backend/open_webui/models/knowledge.py +++ b/backend/open_webui/models/knowledge.py @@ -15,9 +15,10 @@ from open_webui.models.files import ( ) from open_webui.models.groups import Groups from open_webui.models.users import User, UserModel, Users, UserResponse +from open_webui.models.access_grants import AccessGrantModel, AccessGrants -from pydantic import BaseModel, ConfigDict +from pydantic import BaseModel, ConfigDict, Field from sqlalchemy import ( BigInteger, Column, @@ -29,9 +30,6 @@ from sqlalchemy import ( or_, ) -from open_webui.utils.access_control import has_access -from open_webui.utils.db.access_control import has_permission - log = logging.getLogger(__name__) @@ -50,22 +48,6 @@ class Knowledge(Base): description = Column(Text) meta = Column(JSON, nullable=True) - access_control = Column(JSON, nullable=True) # Controls data access levels. - # Defines access control rules for this entry. - # - `None`: Public access, available to all users with the "user" role. - # - `{}`: Private access, restricted exclusively to the owner. - # - Custom permissions: Specific access control for reading and writing; - # Can specify group or user-level restrictions: - # { - # "read": { - # "group_ids": ["group_id1", "group_id2"], - # "user_ids": ["user_id1", "user_id2"] - # }, - # "write": { - # "group_ids": ["group_id1", "group_id2"], - # "user_ids": ["user_id1", "user_id2"] - # } - # } created_at = Column(BigInteger) updated_at = Column(BigInteger) @@ -82,7 +64,7 @@ class KnowledgeModel(BaseModel): meta: Optional[dict] = None - access_control: Optional[dict] = None + access_grants: list[AccessGrantModel] = Field(default_factory=list) created_at: int # timestamp in epoch updated_at: int # timestamp in epoch @@ -139,7 +121,7 @@ class KnowledgeUserResponse(KnowledgeUserModel): class KnowledgeForm(BaseModel): name: str description: str - access_control: Optional[dict] = None + access_grants: Optional[list[dict]] = None class FileUserResponse(FileModelResponse): @@ -157,27 +139,47 @@ class KnowledgeFileListResponse(BaseModel): class KnowledgeTable: + def _get_access_grants( + self, knowledge_id: str, db: Optional[Session] = None + ) -> list[AccessGrantModel]: + return AccessGrants.get_grants_by_resource("knowledge", knowledge_id, db=db) + + def _to_knowledge_model( + self, knowledge: Knowledge, db: Optional[Session] = None + ) -> KnowledgeModel: + knowledge_data = KnowledgeModel.model_validate(knowledge).model_dump( + exclude={"access_grants"} + ) + knowledge_data["access_grants"] = self._get_access_grants( + knowledge_data["id"], db=db + ) + return KnowledgeModel.model_validate(knowledge_data) + def insert_new_knowledge( self, user_id: str, form_data: KnowledgeForm, db: Optional[Session] = None ) -> Optional[KnowledgeModel]: with get_db_context(db) as db: knowledge = KnowledgeModel( **{ - **form_data.model_dump(), + **form_data.model_dump(exclude={"access_grants"}), "id": str(uuid.uuid4()), "user_id": user_id, "created_at": int(time.time()), "updated_at": int(time.time()), + "access_grants": [], } ) try: - result = Knowledge(**knowledge.model_dump()) + result = Knowledge(**knowledge.model_dump(exclude={"access_grants"})) db.add(result) db.commit() db.refresh(result) + AccessGrants.set_access_grants( + "knowledge", result.id, form_data.access_grants, db=db + ) if result: - return KnowledgeModel.model_validate(result) + return self._to_knowledge_model(result, db=db) else: return None except Exception: @@ -201,7 +203,7 @@ class KnowledgeTable: knowledge_bases.append( KnowledgeUserModel.model_validate( { - **KnowledgeModel.model_validate(knowledge).model_dump(), + **self._to_knowledge_model(knowledge, db=db).model_dump(), "user": user.model_dump() if user else None, } ) @@ -241,7 +243,14 @@ class KnowledgeTable: elif view_option == "shared": query = query.filter(Knowledge.user_id != user_id) - query = has_permission(db, Knowledge, query, filter) + query = AccessGrants.has_permission_filter( + db=db, + query=query, + DocumentModel=Knowledge, + filter=filter, + resource_type="knowledge", + permission="read", + ) query = query.order_by(Knowledge.updated_at.desc(), Knowledge.id.asc()) @@ -258,8 +267,8 @@ class KnowledgeTable: knowledge_bases.append( KnowledgeUserModel.model_validate( { - **KnowledgeModel.model_validate( - knowledge_base + **self._to_knowledge_model( + knowledge_base, db=db ).model_dump(), "user": ( UserModel.model_validate(user).model_dump() @@ -294,7 +303,14 @@ class KnowledgeTable: # Apply access-control directly to the joined query # This makes the database handle filtering, even with 10k+ KBs - query = has_permission(db, Knowledge, query, filter) + query = AccessGrants.has_permission_filter( + db=db, + query=query, + DocumentModel=Knowledge, + filter=filter, + resource_type="knowledge", + permission="read", + ) # Apply filename search if filter: @@ -327,8 +343,8 @@ class KnowledgeTable: if user else None ), - collection=KnowledgeModel.model_validate( - knowledge + collection=self._to_knowledge_model( + knowledge, db=db ).model_dump(), ) ) @@ -350,7 +366,14 @@ class KnowledgeTable: user_group_ids = { group.id for group in Groups.get_groups_by_member_id(user_id, db=db) } - return has_access(user_id, permission, knowledge.access_control, user_group_ids) + return AccessGrants.has_access( + user_id=user_id, + resource_type="knowledge", + resource_id=knowledge.id, + permission=permission, + user_group_ids=user_group_ids, + db=db, + ) def get_knowledge_bases_by_user_id( self, user_id: str, permission: str = "write", db: Optional[Session] = None @@ -363,8 +386,13 @@ class KnowledgeTable: knowledge_base for knowledge_base in knowledge_bases if knowledge_base.user_id == user_id - or has_access( - user_id, permission, knowledge_base.access_control, user_group_ids + or AccessGrants.has_access( + user_id=user_id, + resource_type="knowledge", + resource_id=knowledge_base.id, + permission=permission, + user_group_ids=user_group_ids, + db=db, ) ] @@ -374,7 +402,9 @@ class KnowledgeTable: try: with get_db_context(db) as db: knowledge = db.query(Knowledge).filter_by(id=id).first() - return KnowledgeModel.model_validate(knowledge) if knowledge else None + return ( + self._to_knowledge_model(knowledge, db=db) if knowledge else None + ) except Exception: return None @@ -391,7 +421,14 @@ class KnowledgeTable: user_group_ids = { group.id for group in Groups.get_groups_by_member_id(user_id, db=db) } - if has_access(user_id, "write", knowledge.access_control, user_group_ids): + if AccessGrants.has_access( + user_id=user_id, + resource_type="knowledge", + resource_id=knowledge.id, + permission="write", + user_group_ids=user_group_ids, + db=db, + ): return knowledge return None @@ -406,9 +443,7 @@ class KnowledgeTable: .filter(KnowledgeFile.file_id == file_id) .all() ) - return [ - KnowledgeModel.model_validate(knowledge) for knowledge in knowledges - ] + return [self._to_knowledge_model(knowledge, db=db) for knowledge in knowledges] except Exception: return [] @@ -591,11 +626,15 @@ class KnowledgeTable: knowledge = self.get_knowledge_by_id(id=id, db=db) db.query(Knowledge).filter_by(id=id).update( { - **form_data.model_dump(), + **form_data.model_dump(exclude={"access_grants"}), "updated_at": int(time.time()), } ) db.commit() + if form_data.access_grants is not None: + AccessGrants.set_access_grants( + "knowledge", id, form_data.access_grants, db=db + ) return self.get_knowledge_by_id(id=id, db=db) except Exception as e: log.exception(e) @@ -622,6 +661,7 @@ class KnowledgeTable: def delete_knowledge_by_id(self, id: str, db: Optional[Session] = None) -> bool: try: with get_db_context(db) as db: + AccessGrants.revoke_all_access("knowledge", id, db=db) db.query(Knowledge).filter_by(id=id).delete() db.commit() return True @@ -631,6 +671,9 @@ class KnowledgeTable: def delete_all_knowledge(self, db: Optional[Session] = None) -> bool: with get_db_context(db) as db: try: + knowledge_ids = [row[0] for row in db.query(Knowledge.id).all()] + for knowledge_id in knowledge_ids: + AccessGrants.revoke_all_access("knowledge", knowledge_id, db=db) db.query(Knowledge).delete() db.commit() diff --git a/backend/open_webui/models/models.py b/backend/open_webui/models/models.py index 5a59861dd7..d523ae0fc1 100755 --- a/backend/open_webui/models/models.py +++ b/backend/open_webui/models/models.py @@ -7,18 +7,16 @@ from open_webui.internal.db import Base, JSONField, get_db, get_db_context from open_webui.models.groups import Groups from open_webui.models.users import User, UserModel, Users, UserResponse +from open_webui.models.access_grants import AccessGrantModel, AccessGrants -from pydantic import BaseModel, ConfigDict +from pydantic import BaseModel, ConfigDict, Field from sqlalchemy import String, cast, or_, and_, func from sqlalchemy.dialects import postgresql, sqlite from sqlalchemy.dialects.postgresql import JSONB -from sqlalchemy import BigInteger, Column, Text, JSON, Boolean - - -from open_webui.utils.access_control import has_access +from sqlalchemy import BigInteger, Column, Text, Boolean log = logging.getLogger(__name__) @@ -80,23 +78,6 @@ class Model(Base): Holds a JSON encoded blob of metadata, see `ModelMeta`. """ - access_control = Column(JSON, nullable=True) # Controls data access levels. - # Defines access control rules for this entry. - # - `None`: Public access, available to all users with the "user" role. - # - `{}`: Private access, restricted exclusively to the owner. - # - Custom permissions: Specific access control for reading and writing; - # Can specify group or user-level restrictions: - # { - # "read": { - # "group_ids": ["group_id1", "group_id2"], - # "user_ids": ["user_id1", "user_id2"] - # }, - # "write": { - # "group_ids": ["group_id1", "group_id2"], - # "user_ids": ["user_id1", "user_id2"] - # } - # } - is_active = Column(Boolean, default=True) updated_at = Column(BigInteger) @@ -112,7 +93,7 @@ class ModelModel(BaseModel): params: ModelParams meta: ModelMeta - access_control: Optional[dict] = None + access_grants: list[AccessGrantModel] = Field(default_factory=list) is_active: bool updated_at: int # timestamp in epoch @@ -154,31 +135,45 @@ class ModelForm(BaseModel): name: str meta: ModelMeta params: ModelParams - access_control: Optional[dict] = None + access_grants: Optional[list[dict]] = None is_active: bool = True class ModelsTable: + def _get_access_grants( + self, model_id: str, db: Optional[Session] = None + ) -> list[AccessGrantModel]: + return AccessGrants.get_grants_by_resource("model", model_id, db=db) + + def _to_model_model(self, model: Model, db: Optional[Session] = None) -> ModelModel: + model_data = ModelModel.model_validate(model).model_dump( + exclude={"access_grants"} + ) + model_data["access_grants"] = self._get_access_grants(model_data["id"], db=db) + return ModelModel.model_validate(model_data) + def insert_new_model( self, form_data: ModelForm, user_id: str, db: Optional[Session] = None ) -> Optional[ModelModel]: - model = ModelModel( - **{ - **form_data.model_dump(), - "user_id": user_id, - "created_at": int(time.time()), - "updated_at": int(time.time()), - } - ) try: with get_db_context(db) as db: - result = Model(**model.model_dump()) + result = Model( + **{ + **form_data.model_dump(exclude={"access_grants"}), + "user_id": user_id, + "created_at": int(time.time()), + "updated_at": int(time.time()), + } + ) db.add(result) db.commit() db.refresh(result) + AccessGrants.set_access_grants( + "model", result.id, form_data.access_grants, db=db + ) if result: - return ModelModel.model_validate(result) + return self._to_model_model(result, db=db) else: return None except Exception as e: @@ -187,7 +182,7 @@ class ModelsTable: def get_all_models(self, db: Optional[Session] = None) -> list[ModelModel]: with get_db_context(db) as db: - return [ModelModel.model_validate(model) for model in db.query(Model).all()] + return [self._to_model_model(model, db=db) for model in db.query(Model).all()] def get_models(self, db: Optional[Session] = None) -> list[ModelUserResponse]: with get_db_context(db) as db: @@ -204,7 +199,7 @@ class ModelsTable: models.append( ModelUserResponse.model_validate( { - **ModelModel.model_validate(model).model_dump(), + **self._to_model_model(model, db=db).model_dump(), "user": user.model_dump() if user else None, } ) @@ -214,7 +209,7 @@ class ModelsTable: def get_base_models(self, db: Optional[Session] = None) -> list[ModelModel]: with get_db_context(db) as db: return [ - ModelModel.model_validate(model) + self._to_model_model(model, db=db) for model in db.query(Model).filter(Model.base_model_id == None).all() ] @@ -229,50 +224,25 @@ class ModelsTable: model for model in models if model.user_id == user_id - or has_access(user_id, permission, model.access_control, user_group_ids) + or AccessGrants.has_access( + user_id=user_id, + resource_type="model", + resource_id=model.id, + permission=permission, + user_group_ids=user_group_ids, + db=db, + ) ] def _has_permission(self, db, query, filter: dict, permission: str = "read"): - group_ids = filter.get("group_ids", []) - user_id = filter.get("user_id") - - dialect_name = db.bind.dialect.name - - # Public access - conditions = [] - if group_ids or user_id: - conditions.extend( - [ - Model.access_control.is_(None), - cast(Model.access_control, String) == "null", - ] - ) - - # User-level permission - if user_id: - conditions.append(Model.user_id == user_id) - - # Group-level permission - if group_ids: - group_conditions = [] - for gid in group_ids: - if dialect_name == "sqlite": - group_conditions.append( - Model.access_control[permission]["group_ids"].contains([gid]) - ) - elif dialect_name == "postgresql": - group_conditions.append( - cast( - Model.access_control[permission]["group_ids"], - JSONB, - ).contains([gid]) - ) - conditions.append(or_(*group_conditions)) - - if conditions: - query = query.filter(or_(*conditions)) - - return query + return AccessGrants.has_permission_filter( + db=db, + query=query, + DocumentModel=Model, + filter=filter, + resource_type="model", + permission=permission, + ) def search_models( self, @@ -358,7 +328,7 @@ class ModelsTable: for model, user in items: models.append( ModelUserResponse( - **ModelModel.model_validate(model).model_dump(), + **self._to_model_model(model, db=db).model_dump(), user=( UserResponse(**UserModel.model_validate(user).model_dump()) if user @@ -375,7 +345,7 @@ class ModelsTable: try: with get_db_context(db) as db: model = db.get(Model, id) - return ModelModel.model_validate(model) + return self._to_model_model(model, db=db) if model else None except Exception: return None @@ -385,7 +355,7 @@ class ModelsTable: try: with get_db_context(db) as db: models = db.query(Model).filter(Model.id.in_(ids)).all() - return [ModelModel.model_validate(model) for model in models] + return [self._to_model_model(model, db=db) for model in models] except Exception: return [] @@ -403,7 +373,7 @@ class ModelsTable: db.commit() db.refresh(model) - return ModelModel.model_validate(model) + return self._to_model_model(model, db=db) except Exception: return None @@ -413,14 +383,16 @@ class ModelsTable: try: with get_db_context(db) as db: # update only the fields that are present in the model - data = model.model_dump(exclude={"id"}) + data = model.model_dump(exclude={"id", "access_grants"}) result = db.query(Model).filter_by(id=id).update(data) db.commit() + if model.access_grants is not None: + AccessGrants.set_access_grants( + "model", id, model.access_grants, db=db + ) - model = db.get(Model, id) - db.refresh(model) - return ModelModel.model_validate(model) + return self.get_model_by_id(id, db=db) except Exception as e: log.exception(f"Failed to update the model by id {id}: {e}") return None @@ -428,6 +400,7 @@ class ModelsTable: def delete_model_by_id(self, id: str, db: Optional[Session] = None) -> bool: try: with get_db_context(db) as db: + AccessGrants.revoke_all_access("model", id, db=db) db.query(Model).filter_by(id=id).delete() db.commit() @@ -438,6 +411,9 @@ class ModelsTable: def delete_all_models(self, db: Optional[Session] = None) -> bool: try: with get_db_context(db) as db: + model_ids = [row[0] for row in db.query(Model.id).all()] + for model_id in model_ids: + AccessGrants.revoke_all_access("model", model_id, db=db) db.query(Model).delete() db.commit() @@ -462,7 +438,7 @@ class ModelsTable: if model.id in existing_ids: db.query(Model).filter_by(id=model.id).update( { - **model.model_dump(), + **model.model_dump(exclude={"access_grants"}), "user_id": user_id, "updated_at": int(time.time()), } @@ -470,22 +446,27 @@ class ModelsTable: else: new_model = Model( **{ - **model.model_dump(), + **model.model_dump(exclude={"access_grants"}), "user_id": user_id, "updated_at": int(time.time()), } ) db.add(new_model) + AccessGrants.set_access_grants( + "model", model.id, model.access_grants, db=db + ) # Remove models that are no longer present for model in existing_models: if model.id not in new_model_ids: + AccessGrants.revoke_all_access("model", model.id, db=db) db.delete(model) db.commit() return [ - ModelModel.model_validate(model) for model in db.query(Model).all() + self._to_model_model(model, db=db) + for model in db.query(Model).all() ] except Exception as e: log.exception(f"Error syncing models for user {user_id}: {e}") diff --git a/backend/open_webui/models/notes.py b/backend/open_webui/models/notes.py index bd23530785..d17c749d1c 100644 --- a/backend/open_webui/models/notes.py +++ b/backend/open_webui/models/notes.py @@ -7,17 +7,13 @@ from functools import lru_cache from sqlalchemy.orm import Session from open_webui.internal.db import Base, get_db, get_db_context from open_webui.models.groups import Groups -from open_webui.utils.access_control import has_access from open_webui.models.users import User, UserModel, Users, UserResponse +from open_webui.models.access_grants import AccessGrantModel, AccessGrants -from pydantic import BaseModel, ConfigDict -from sqlalchemy import BigInteger, Boolean, Column, String, Text, JSON -from sqlalchemy.dialects.postgresql import JSONB - - -from sqlalchemy import or_, func, select, and_, text, cast, or_, and_, func -from sqlalchemy.sql import exists +from pydantic import BaseModel, ConfigDict, Field +from sqlalchemy import BigInteger, Column, Text, JSON +from sqlalchemy import or_, func, cast #################### # Note DB Schema @@ -34,8 +30,6 @@ class Note(Base): data = Column(JSON, nullable=True) meta = Column(JSON, nullable=True) - access_control = Column(JSON, nullable=True) - created_at = Column(BigInteger) updated_at = Column(BigInteger) @@ -50,7 +44,7 @@ class NoteModel(BaseModel): data: Optional[dict] = None meta: Optional[dict] = None - access_control: Optional[dict] = None + access_grants: list[AccessGrantModel] = Field(default_factory=list) created_at: int # timestamp in epoch updated_at: int # timestamp in epoch @@ -65,14 +59,14 @@ class NoteForm(BaseModel): title: str data: Optional[dict] = None meta: Optional[dict] = None - access_control: Optional[dict] = None + access_grants: Optional[list[dict]] = None class NoteUpdateForm(BaseModel): title: Optional[str] = None data: Optional[dict] = None meta: Optional[dict] = None - access_control: Optional[dict] = None + access_grants: Optional[list[dict]] = None class NoteUserResponse(NoteModel): @@ -94,122 +88,25 @@ class NoteListResponse(BaseModel): class NoteTable: + def _get_access_grants( + self, note_id: str, db: Optional[Session] = None + ) -> list[AccessGrantModel]: + return AccessGrants.get_grants_by_resource("note", note_id, db=db) + + def _to_note_model(self, note: Note, db: Optional[Session] = None) -> NoteModel: + note_data = NoteModel.model_validate(note).model_dump(exclude={"access_grants"}) + note_data["access_grants"] = self._get_access_grants(note_data["id"], db=db) + return NoteModel.model_validate(note_data) + def _has_permission(self, db, query, filter: dict, permission: str = "read"): - group_ids = filter.get("group_ids", []) - user_id = filter.get("user_id") - dialect_name = db.bind.dialect.name - - conditions = [] - - # Handle read_only permission separately - if permission == "read_only": - # For read_only, we want items where: - # 1. User has explicit read permission (via groups or user-level) - # 2. BUT does NOT have write permission - # 3. Public items are NOT considered read_only - - read_conditions = [] - - # Group-level read permission - if group_ids: - group_read_conditions = [] - for gid in group_ids: - if dialect_name == "sqlite": - group_read_conditions.append( - Note.access_control["read"]["group_ids"].contains([gid]) - ) - elif dialect_name == "postgresql": - group_read_conditions.append( - cast( - Note.access_control["read"]["group_ids"], - JSONB, - ).contains([gid]) - ) - - if group_read_conditions: - read_conditions.append(or_(*group_read_conditions)) - - # Combine read conditions - if read_conditions: - has_read = or_(*read_conditions) - else: - # If no read conditions, return empty result - return query.filter(False) - - # Now exclude items where user has write permission - write_exclusions = [] - - # Exclude items owned by user (they have implicit write) - if user_id: - write_exclusions.append(Note.user_id != user_id) - - # Exclude items where user has explicit write permission via groups - if group_ids: - group_write_conditions = [] - for gid in group_ids: - if dialect_name == "sqlite": - group_write_conditions.append( - Note.access_control["write"]["group_ids"].contains([gid]) - ) - elif dialect_name == "postgresql": - group_write_conditions.append( - cast( - Note.access_control["write"]["group_ids"], - JSONB, - ).contains([gid]) - ) - - if group_write_conditions: - # User should NOT have write permission - write_exclusions.append(~or_(*group_write_conditions)) - - # Exclude public items (items without access_control) - write_exclusions.append(Note.access_control.isnot(None)) - write_exclusions.append(cast(Note.access_control, String) != "null") - - # Combine: has read AND does not have write AND not public - if write_exclusions: - query = query.filter(and_(has_read, *write_exclusions)) - else: - query = query.filter(has_read) - - return query - - # Original logic for other permissions (read, write, etc.) - # Public access conditions - if group_ids or user_id: - conditions.extend( - [ - Note.access_control.is_(None), - cast(Note.access_control, String) == "null", - ] - ) - - # User-level permission (owner has all permissions) - if user_id: - conditions.append(Note.user_id == user_id) - - # Group-level permission - if group_ids: - group_conditions = [] - for gid in group_ids: - if dialect_name == "sqlite": - group_conditions.append( - Note.access_control[permission]["group_ids"].contains([gid]) - ) - elif dialect_name == "postgresql": - group_conditions.append( - cast( - Note.access_control[permission]["group_ids"], - JSONB, - ).contains([gid]) - ) - conditions.append(or_(*group_conditions)) - - if conditions: - query = query.filter(or_(*conditions)) - - return query + return AccessGrants.has_permission_filter( + db=db, + query=query, + DocumentModel=Note, + filter=filter, + resource_type="note", + permission=permission, + ) def insert_new_note( self, user_id: str, form_data: NoteForm, db: Optional[Session] = None @@ -219,17 +116,21 @@ class NoteTable: **{ "id": str(uuid.uuid4()), "user_id": user_id, - **form_data.model_dump(), + **form_data.model_dump(exclude={"access_grants"}), "created_at": int(time.time_ns()), "updated_at": int(time.time_ns()), + "access_grants": [], } ) - new_note = Note(**note.model_dump()) + new_note = Note(**note.model_dump(exclude={"access_grants"})) db.add(new_note) db.commit() - return note + AccessGrants.set_access_grants( + "note", note.id, form_data.access_grants, db=db + ) + return self._to_note_model(new_note, db=db) def get_notes( self, skip: int = 0, limit: int = 50, db: Optional[Session] = None @@ -241,7 +142,7 @@ class NoteTable: if limit is not None: query = query.limit(limit) notes = query.all() - return [NoteModel.model_validate(note) for note in notes] + return [self._to_note_model(note, db=db) for note in notes] def search_notes( self, @@ -330,7 +231,7 @@ class NoteTable: for note, user in items: notes.append( NoteUserResponse( - **NoteModel.model_validate(note).model_dump(), + **self._to_note_model(note, db=db).model_dump(), user=( UserResponse(**UserModel.model_validate(user).model_dump()) if user @@ -365,14 +266,14 @@ class NoteTable: query = query.limit(limit) notes = query.all() - return [NoteModel.model_validate(note) for note in notes] + return [self._to_note_model(note, db=db) for note in notes] def get_note_by_id( self, id: str, db: Optional[Session] = None ) -> Optional[NoteModel]: with get_db_context(db) as db: note = db.query(Note).filter(Note.id == id).first() - return NoteModel.model_validate(note) if note else None + return self._to_note_model(note, db=db) if note else None def update_note_by_id( self, id: str, form_data: NoteUpdateForm, db: Optional[Session] = None @@ -391,17 +292,20 @@ class NoteTable: if "meta" in form_data: note.meta = {**note.meta, **form_data["meta"]} - if "access_control" in form_data: - note.access_control = form_data["access_control"] + if "access_grants" in form_data: + AccessGrants.set_access_grants( + "note", id, form_data["access_grants"], db=db + ) note.updated_at = int(time.time_ns()) db.commit() - return NoteModel.model_validate(note) if note else None + return self._to_note_model(note, db=db) if note else None def delete_note_by_id(self, id: str, db: Optional[Session] = None) -> bool: try: with get_db_context(db) as db: + AccessGrants.revoke_all_access("note", id, db=db) db.query(Note).filter(Note.id == id).delete() db.commit() return True diff --git a/backend/open_webui/models/prompt_history.py b/backend/open_webui/models/prompt_history.py index ea7f566fb1..0f5e7cea87 100644 --- a/backend/open_webui/models/prompt_history.py +++ b/backend/open_webui/models/prompt_history.py @@ -45,6 +45,7 @@ class PromptHistoryModel(BaseModel): class PromptHistoryResponse(PromptHistoryModel): """Response model with user info.""" + user: Optional[UserResponse] = None @@ -91,16 +92,20 @@ class PromptHistoryTable: .limit(limit) .all() ) - + # Get user info for each entry user_ids = list(set(e.user_id for e in entries)) users = Users.get_users_by_user_ids(user_ids, db=db) if user_ids else [] users_dict = {user.id: user for user in users} - + return [ PromptHistoryResponse( **PromptHistoryModel.model_validate(entry).model_dump(), - user=users_dict.get(entry.user_id).model_dump() if users_dict.get(entry.user_id) else None, + user=( + users_dict.get(entry.user_id).model_dump() + if users_dict.get(entry.user_id) + else None + ), ) for entry in entries ] @@ -112,7 +117,9 @@ class PromptHistoryTable: ) -> Optional[PromptHistoryModel]: """Get a specific history entry by ID.""" with get_db_context(db) as db: - entry = db.query(PromptHistory).filter(PromptHistory.id == history_id).first() + entry = ( + db.query(PromptHistory).filter(PromptHistory.id == history_id).first() + ) if entry: return PromptHistoryModel.model_validate(entry) return None @@ -155,27 +162,31 @@ class PromptHistoryTable: ) -> Optional[dict]: """Compute diff between two history entries.""" with get_db_context(db) as db: - from_entry = db.query(PromptHistory).filter(PromptHistory.id == from_id).first() + from_entry = ( + db.query(PromptHistory).filter(PromptHistory.id == from_id).first() + ) to_entry = db.query(PromptHistory).filter(PromptHistory.id == to_id).first() - + if not from_entry or not to_entry: return None - + from_snapshot = from_entry.snapshot to_snapshot = to_entry.snapshot - + # Compute diff for content field from_content = from_snapshot.get("content", "") to_content = to_snapshot.get("content", "") - - diff_lines = list(difflib.unified_diff( - from_content.splitlines(keepends=True), - to_content.splitlines(keepends=True), - fromfile=f"v{from_id[:8]}", - tofile=f"v{to_id[:8]}", - lineterm="", - )) - + + diff_lines = list( + difflib.unified_diff( + from_content.splitlines(keepends=True), + to_content.splitlines(keepends=True), + fromfile=f"v{from_id[:8]}", + tofile=f"v{to_id[:8]}", + lineterm="", + ) + ) + return { "from_id": from_id, "to_id": to_id, @@ -183,7 +194,6 @@ class PromptHistoryTable: "to_snapshot": to_snapshot, "content_diff": diff_lines, "name_changed": from_snapshot.get("name") != to_snapshot.get("name"), - "access_control_changed": from_snapshot.get("access_control") != to_snapshot.get("access_control"), } def delete_history_by_prompt_id( @@ -193,7 +203,9 @@ class PromptHistoryTable: ) -> bool: """Delete all history entries for a prompt.""" with get_db_context(db) as db: - db.query(PromptHistory).filter(PromptHistory.prompt_id == prompt_id).delete() + db.query(PromptHistory).filter( + PromptHistory.prompt_id == prompt_id + ).delete() db.commit() return True diff --git a/backend/open_webui/models/prompts.py b/backend/open_webui/models/prompts.py index 4a85ba9029..544aea767b 100644 --- a/backend/open_webui/models/prompts.py +++ b/backend/open_webui/models/prompts.py @@ -7,15 +7,13 @@ from open_webui.internal.db import Base, JSONField, get_db, get_db_context from open_webui.models.groups import Groups from open_webui.models.users import Users, UserResponse from open_webui.models.prompt_history import PromptHistories +from open_webui.models.access_grants import AccessGrantModel, AccessGrants -from pydantic import BaseModel, ConfigDict +from pydantic import BaseModel, ConfigDict, Field from sqlalchemy import BigInteger, Boolean, Column, String, Text, JSON, or_, func, cast -from open_webui.utils.access_control import has_access - - #################### # Prompts DB Schema #################### @@ -37,23 +35,6 @@ class Prompt(Base): created_at = Column(BigInteger, nullable=True) updated_at = Column(BigInteger, nullable=True) - access_control = Column(JSON, nullable=True) # Controls data access levels. - # Defines access control rules for this entry. - # - `None`: Public access, available to all users with the "user" role. - # - `{}`: Private access, restricted exclusively to the owner. - # - Custom permissions: Specific access control for reading and writing; - # Can specify group or user-level restrictions: - # { - # "read": { - # "group_ids": ["group_id1", "group_id2"], - # "user_ids": ["user_id1", "user_id2"] - # }, - # "write": { - # "group_ids": ["group_id1", "group_id2"], - # "user_ids": ["user_id1", "user_id2"] - # } - # } - class PromptModel(BaseModel): id: Optional[str] = None @@ -68,7 +49,7 @@ class PromptModel(BaseModel): version_id: Optional[str] = None created_at: Optional[int] = None updated_at: Optional[int] = None - access_control: Optional[dict] = None + access_grants: list[AccessGrantModel] = Field(default_factory=list) model_config = ConfigDict(from_attributes=True) @@ -104,13 +85,27 @@ class PromptForm(BaseModel): data: Optional[dict] = None meta: Optional[dict] = None tags: Optional[list[str]] = None - access_control: Optional[dict] = None + access_grants: Optional[list[dict]] = None version_id: Optional[str] = None # Active version commit_message: Optional[str] = None # For history tracking is_production: Optional[bool] = True # Whether to set new version as production class PromptsTable: + def _get_access_grants( + self, prompt_id: str, db: Optional[Session] = None + ) -> list[AccessGrantModel]: + return AccessGrants.get_grants_by_resource("prompt", prompt_id, db=db) + + def _to_prompt_model( + self, prompt: Prompt, db: Optional[Session] = None + ) -> PromptModel: + prompt_data = PromptModel.model_validate(prompt).model_dump( + exclude={"access_grants"} + ) + prompt_data["access_grants"] = self._get_access_grants(prompt_data["id"], db=db) + return PromptModel.model_validate(prompt_data) + def insert_new_prompt( self, user_id: str, form_data: PromptForm, db: Optional[Session] = None ) -> Optional[PromptModel]: @@ -126,7 +121,7 @@ class PromptsTable: data=form_data.data or {}, meta=form_data.meta or {}, tags=form_data.tags or [], - access_control=form_data.access_control, + access_grants=[], is_active=True, created_at=now, updated_at=now, @@ -134,12 +129,16 @@ class PromptsTable: try: with get_db_context(db) as db: - result = Prompt(**prompt.model_dump()) + result = Prompt(**prompt.model_dump(exclude={"access_grants"})) db.add(result) db.commit() db.refresh(result) + AccessGrants.set_access_grants( + "prompt", prompt_id, form_data.access_grants, db=db + ) if result: + current_access_grants = self._get_access_grants(prompt_id, db=db) snapshot = { "name": form_data.name, "content": form_data.content, @@ -147,7 +146,7 @@ class PromptsTable: "data": form_data.data or {}, "meta": form_data.meta or {}, "tags": form_data.tags or [], - "access_control": form_data.access_control, + "access_grants": [grant.model_dump() for grant in current_access_grants], } history_entry = PromptHistories.create_history_entry( @@ -165,7 +164,7 @@ class PromptsTable: db.commit() db.refresh(result) - return PromptModel.model_validate(result) + return self._to_prompt_model(result, db=db) else: return None except Exception: @@ -179,7 +178,7 @@ class PromptsTable: with get_db_context(db) as db: prompt = db.query(Prompt).filter_by(id=prompt_id).first() if prompt: - return PromptModel.model_validate(prompt) + return self._to_prompt_model(prompt, db=db) return None except Exception: return None @@ -191,7 +190,7 @@ class PromptsTable: with get_db_context(db) as db: prompt = db.query(Prompt).filter_by(command=command).first() if prompt: - return PromptModel.model_validate(prompt) + return self._to_prompt_model(prompt, db=db) return None except Exception: return None @@ -216,7 +215,7 @@ class PromptsTable: prompts.append( PromptUserResponse.model_validate( { - **PromptModel.model_validate(prompt).model_dump(), + **self._to_prompt_model(prompt, db=db).model_dump(), "user": user.model_dump() if user else None, } ) @@ -236,7 +235,14 @@ class PromptsTable: prompt for prompt in prompts if prompt.user_id == user_id - or has_access(user_id, permission, prompt.access_control, user_group_ids) + or AccessGrants.has_access( + user_id=user_id, + resource_type="prompt", + resource_id=prompt.id, + permission=permission, + user_group_ids=user_group_ids, + db=db, + ) ] def search_prompts( @@ -273,17 +279,15 @@ class PromptsTable: elif view_option == "shared": query = query.filter(Prompt.user_id != user_id) - # Apply access control filtering - group_ids = filter.get("group_ids", []) - filter_user_id = filter.get("user_id") - - if filter_user_id: - # User must have access: owner OR public OR explicit access - access_conditions = [ - Prompt.user_id == filter_user_id, # Owner - Prompt.access_control == None, # Public - ] - query = query.filter(or_(*access_conditions)) + # Apply access grant filtering + query = AccessGrants.has_permission_filter( + db=db, + query=query, + DocumentModel=Prompt, + filter=filter, + resource_type="prompt", + permission="read", + ) tag = filter.get("tag") if tag: @@ -329,7 +333,7 @@ class PromptsTable: for prompt, user in items: prompts.append( PromptUserResponse( - **PromptModel.model_validate(prompt).model_dump(), + **self._to_prompt_model(prompt, db=db).model_dump(), user=( UserResponse(**UserModel.model_validate(user).model_dump()) if user @@ -358,12 +362,13 @@ class PromptsTable: prompt.id, db=db ) parent_id = latest_history.id if latest_history else None + current_access_grants = self._get_access_grants(prompt.id, db=db) # Check if content changed to decide on history creation content_changed = ( prompt.name != form_data.name or prompt.content != form_data.content - or prompt.access_control != form_data.access_control + or form_data.access_grants is not None ) # Update prompt fields @@ -371,8 +376,12 @@ class PromptsTable: prompt.content = form_data.content prompt.data = form_data.data or prompt.data prompt.meta = form_data.meta or prompt.meta - prompt.access_control = form_data.access_control prompt.updated_at = int(time.time()) + if form_data.access_grants is not None: + AccessGrants.set_access_grants( + "prompt", prompt.id, form_data.access_grants, db=db + ) + current_access_grants = self._get_access_grants(prompt.id, db=db) db.commit() @@ -384,7 +393,9 @@ class PromptsTable: "command": command, "data": form_data.data or {}, "meta": form_data.meta or {}, - "access_control": form_data.access_control, + "access_grants": [ + grant.model_dump() for grant in current_access_grants + ], } history_entry = PromptHistories.create_history_entry( @@ -401,7 +412,7 @@ class PromptsTable: prompt.version_id = history_entry.id db.commit() - return PromptModel.model_validate(prompt) + return self._to_prompt_model(prompt, db=db) except Exception: return None @@ -422,13 +433,14 @@ class PromptsTable: prompt.id, db=db ) parent_id = latest_history.id if latest_history else None + current_access_grants = self._get_access_grants(prompt.id, db=db) # Check if content changed to decide on history creation content_changed = ( prompt.name != form_data.name or prompt.command != form_data.command or prompt.content != form_data.content - or prompt.access_control != form_data.access_control + or form_data.access_grants is not None or (form_data.tags is not None and prompt.tags != form_data.tags) ) @@ -438,10 +450,15 @@ class PromptsTable: prompt.content = form_data.content prompt.data = form_data.data or prompt.data prompt.meta = form_data.meta or prompt.meta - prompt.access_control = form_data.access_control if form_data.tags is not None: prompt.tags = form_data.tags + + if form_data.access_grants is not None: + AccessGrants.set_access_grants( + "prompt", prompt.id, form_data.access_grants, db=db + ) + current_access_grants = self._get_access_grants(prompt.id, db=db) prompt.updated_at = int(time.time()) @@ -456,7 +473,9 @@ class PromptsTable: "data": form_data.data or {}, "meta": form_data.meta or {}, "tags": prompt.tags or [], - "access_control": form_data.access_control, + "access_grants": [ + grant.model_dump() for grant in current_access_grants + ], } history_entry = PromptHistories.create_history_entry( @@ -473,7 +492,7 @@ class PromptsTable: prompt.version_id = history_entry.id db.commit() - return PromptModel.model_validate(prompt) + return self._to_prompt_model(prompt, db=db) except Exception: return None @@ -501,7 +520,7 @@ class PromptsTable: prompt.updated_at = int(time.time()) db.commit() - return PromptModel.model_validate(prompt) + return self._to_prompt_model(prompt, db=db) except Exception: return None @@ -533,13 +552,13 @@ class PromptsTable: prompt.data = snapshot.get("data", prompt.data) prompt.meta = snapshot.get("meta", prompt.meta) prompt.tags = snapshot.get("tags", prompt.tags) - # Note: command and access_control are not restored from snapshot + # Note: command and access_grants are not restored from snapshot prompt.version_id = version_id prompt.updated_at = int(time.time()) db.commit() - return PromptModel.model_validate(prompt) + return self._to_prompt_model(prompt, db=db) except Exception: return None @@ -552,6 +571,7 @@ class PromptsTable: prompt = db.query(Prompt).filter_by(command=command).first() if prompt: PromptHistories.delete_history_by_prompt_id(prompt.id, db=db) + AccessGrants.revoke_all_access("prompt", prompt.id, db=db) prompt.is_active = False prompt.updated_at = int(time.time()) @@ -568,6 +588,7 @@ class PromptsTable: prompt = db.query(Prompt).filter_by(id=prompt_id).first() if prompt: PromptHistories.delete_history_by_prompt_id(prompt.id, db=db) + AccessGrants.revoke_all_access("prompt", prompt.id, db=db) prompt.is_active = False prompt.updated_at = int(time.time()) @@ -586,6 +607,7 @@ class PromptsTable: prompt = db.query(Prompt).filter_by(command=command).first() if prompt: PromptHistories.delete_history_by_prompt_id(prompt.id, db=db) + AccessGrants.revoke_all_access("prompt", prompt.id, db=db) # Delete prompt db.query(Prompt).filter_by(command=command).delete() diff --git a/backend/open_webui/models/tools.py b/backend/open_webui/models/tools.py index cd7d0bd1a0..da439161e9 100644 --- a/backend/open_webui/models/tools.py +++ b/backend/open_webui/models/tools.py @@ -6,11 +6,10 @@ from sqlalchemy.orm import Session from open_webui.internal.db import Base, JSONField, get_db, get_db_context from open_webui.models.users import Users, UserResponse from open_webui.models.groups import Groups +from open_webui.models.access_grants import AccessGrantModel, AccessGrants -from pydantic import BaseModel, ConfigDict -from sqlalchemy import BigInteger, Column, String, Text, JSON - -from open_webui.utils.access_control import has_access +from pydantic import BaseModel, ConfigDict, Field +from sqlalchemy import BigInteger, Column, String, Text log = logging.getLogger(__name__) @@ -31,23 +30,6 @@ class Tool(Base): meta = Column(JSONField) valves = Column(JSONField) - access_control = Column(JSON, nullable=True) # Controls data access levels. - # Defines access control rules for this entry. - # - `None`: Public access, available to all users with the "user" role. - # - `{}`: Private access, restricted exclusively to the owner. - # - Custom permissions: Specific access control for reading and writing; - # Can specify group or user-level restrictions: - # { - # "read": { - # "group_ids": ["group_id1", "group_id2"], - # "user_ids": ["user_id1", "user_id2"] - # }, - # "write": { - # "group_ids": ["group_id1", "group_id2"], - # "user_ids": ["user_id1", "user_id2"] - # } - # } - updated_at = Column(BigInteger) created_at = Column(BigInteger) @@ -64,7 +46,7 @@ class ToolModel(BaseModel): content: str specs: list[dict] meta: ToolMeta - access_control: Optional[dict] = None + access_grants: list[AccessGrantModel] = Field(default_factory=list) updated_at: int # timestamp in epoch created_at: int # timestamp in epoch @@ -86,7 +68,7 @@ class ToolResponse(BaseModel): user_id: str name: str meta: ToolMeta - access_control: Optional[dict] = None + access_grants: list[AccessGrantModel] = Field(default_factory=list) updated_at: int # timestamp in epoch created_at: int # timestamp in epoch @@ -106,7 +88,7 @@ class ToolForm(BaseModel): name: str content: str meta: ToolMeta - access_control: Optional[dict] = None + access_grants: Optional[list[dict]] = None class ToolValves(BaseModel): @@ -114,6 +96,16 @@ class ToolValves(BaseModel): class ToolsTable: + def _get_access_grants( + self, tool_id: str, db: Optional[Session] = None + ) -> list[AccessGrantModel]: + return AccessGrants.get_grants_by_resource("tool", tool_id, db=db) + + def _to_tool_model(self, tool: Tool, db: Optional[Session] = None) -> ToolModel: + tool_data = ToolModel.model_validate(tool).model_dump(exclude={"access_grants"}) + tool_data["access_grants"] = self._get_access_grants(tool_data["id"], db=db) + return ToolModel.model_validate(tool_data) + def insert_new_tool( self, user_id: str, @@ -122,23 +114,24 @@ class ToolsTable: db: Optional[Session] = None, ) -> Optional[ToolModel]: with get_db_context(db) as db: - tool = ToolModel( - **{ - **form_data.model_dump(), - "specs": specs, - "user_id": user_id, - "updated_at": int(time.time()), - "created_at": int(time.time()), - } - ) - try: - result = Tool(**tool.model_dump()) + result = Tool( + **{ + **form_data.model_dump(exclude={"access_grants"}), + "specs": specs, + "user_id": user_id, + "updated_at": int(time.time()), + "created_at": int(time.time()), + } + ) db.add(result) db.commit() db.refresh(result) + AccessGrants.set_access_grants( + "tool", result.id, form_data.access_grants, db=db + ) if result: - return ToolModel.model_validate(result) + return self._to_tool_model(result, db=db) else: return None except Exception as e: @@ -151,7 +144,7 @@ class ToolsTable: try: with get_db_context(db) as db: tool = db.get(Tool, id) - return ToolModel.model_validate(tool) + return self._to_tool_model(tool, db=db) if tool else None except Exception: return None @@ -170,7 +163,7 @@ class ToolsTable: tools.append( ToolUserModel.model_validate( { - **ToolModel.model_validate(tool).model_dump(), + **self._to_tool_model(tool, db=db).model_dump(), "user": user.model_dump() if user else None, } ) @@ -189,7 +182,14 @@ class ToolsTable: tool for tool in tools if tool.user_id == user_id - or has_access(user_id, permission, tool.access_control, user_group_ids) + or AccessGrants.has_access( + user_id=user_id, + resource_type="tool", + resource_id=tool.id, + permission=permission, + user_group_ids=user_group_ids, + db=db, + ) ] def get_tool_valves_by_id( @@ -266,20 +266,24 @@ class ToolsTable: ) -> Optional[ToolModel]: try: with get_db_context(db) as db: + access_grants = updated.pop("access_grants", None) db.query(Tool).filter_by(id=id).update( {**updated, "updated_at": int(time.time())} ) db.commit() + if access_grants is not None: + AccessGrants.set_access_grants("tool", id, access_grants, db=db) tool = db.query(Tool).get(id) db.refresh(tool) - return ToolModel.model_validate(tool) + return self._to_tool_model(tool, db=db) except Exception: return None def delete_tool_by_id(self, id: str, db: Optional[Session] = None) -> bool: try: with get_db_context(db) as db: + AccessGrants.revoke_all_access("tool", id, db=db) db.query(Tool).filter_by(id=id).delete() db.commit() diff --git a/backend/open_webui/models/users.py b/backend/open_webui/models/users.py index 43c5a94f81..4a526d03f3 100644 --- a/backend/open_webui/models/users.py +++ b/backend/open_webui/models/users.py @@ -184,6 +184,9 @@ class UserInfoResponse(UserStatus): name: str email: str role: str + bio: Optional[str] = None + groups: Optional[list] = [] + is_active: bool = False class UserIdNameResponse(BaseModel): diff --git a/backend/open_webui/retrieval/utils.py b/backend/open_webui/retrieval/utils.py index 56315c73fd..e2e1b85770 100644 --- a/backend/open_webui/retrieval/utils.py +++ b/backend/open_webui/retrieval/utils.py @@ -29,9 +29,9 @@ from open_webui.models.knowledge import Knowledges from open_webui.models.chats import Chats from open_webui.models.notes import Notes +from open_webui.models.access_grants import AccessGrants from open_webui.retrieval.vector.main import GetResult -from open_webui.utils.access_control import has_access from open_webui.utils.headers import include_user_info_headers from open_webui.utils.misc import get_message_list @@ -999,7 +999,12 @@ async def get_sources_from_items( if note and ( user.role == "admin" or note.user_id == user.id - or has_access(user.id, "read", note.access_control) + or AccessGrants.has_access( + user_id=user.id, + resource_type="note", + resource_id=note.id, + permission="read", + ) ): # User has access to the note query_result = { @@ -1091,7 +1096,12 @@ async def get_sources_from_items( if knowledge_base and ( user.role == "admin" or knowledge_base.user_id == user.id - or has_access(user.id, "read", knowledge_base.access_control) + or AccessGrants.has_access( + user_id=user.id, + resource_type="knowledge", + resource_id=knowledge_base.id, + permission="read", + ) ): if ( item.get("context") == "full" @@ -1100,7 +1110,12 @@ async def get_sources_from_items( if knowledge_base and ( user.role == "admin" or knowledge_base.user_id == user.id - or has_access(user.id, "read", knowledge_base.access_control) + or AccessGrants.has_access( + user_id=user.id, + resource_type="knowledge", + resource_id=knowledge_base.id, + permission="read", + ) ): files = Knowledges.get_files_by_id(knowledge_base.id) diff --git a/backend/open_webui/retrieval/vector/dbs/opensearch.py b/backend/open_webui/retrieval/vector/dbs/opensearch.py index dc9c35805e..ed5a931c68 100644 --- a/backend/open_webui/retrieval/vector/dbs/opensearch.py +++ b/backend/open_webui/retrieval/vector/dbs/opensearch.py @@ -211,7 +211,7 @@ class OpenSearchClient(VectorDBBase): for item in batch ] bulk(self.client, actions) - self.client.indices.refresh(self._get_index_name(collection_name)) + self.client.indices.refresh(index=self._get_index_name(collection_name)) def upsert(self, collection_name: str, items: list[VectorItem]): self._create_index_if_not_exists( @@ -234,7 +234,7 @@ class OpenSearchClient(VectorDBBase): for item in batch ] bulk(self.client, actions) - self.client.indices.refresh(self._get_index_name(collection_name)) + self.client.indices.refresh(index=self._get_index_name(collection_name)) def delete( self, @@ -263,7 +263,7 @@ class OpenSearchClient(VectorDBBase): self.client.delete_by_query( index=self._get_index_name(collection_name), body=query_body ) - self.client.indices.refresh(self._get_index_name(collection_name)) + self.client.indices.refresh(index=self._get_index_name(collection_name)) def reset(self): indices = self.client.indices.get(index=f"{self.index_prefix}_*") diff --git a/backend/open_webui/retrieval/web/external.py b/backend/open_webui/retrieval/web/external.py index 527c918a47..1dd63273ae 100644 --- a/backend/open_webui/retrieval/web/external.py +++ b/backend/open_webui/retrieval/web/external.py @@ -8,6 +8,7 @@ from fastapi import Request from open_webui.retrieval.web.main import SearchResult, get_filtered_results from open_webui.utils.headers import include_user_info_headers +from open_webui.env import FORWARD_SESSION_INFO_HEADER_CHAT_ID log = logging.getLogger(__name__) @@ -31,7 +32,7 @@ def search_external( chat_id = getattr(request.state, "chat_id", None) if chat_id: - headers["X-OpenWebUI-Chat-Id"] = str(chat_id) + headers[FORWARD_SESSION_INFO_HEADER_CHAT_ID] = str(chat_id) response = requests.post( external_url, diff --git a/backend/open_webui/retrieval/web/utils.py b/backend/open_webui/retrieval/web/utils.py index 6c1ea4b1bf..45787cb4bd 100644 --- a/backend/open_webui/retrieval/web/utils.py +++ b/backend/open_webui/retrieval/web/utils.py @@ -174,7 +174,7 @@ class URLProcessingMixin: def _safe_process_url_sync(self, url: str) -> bool: """Synchronous version of safety checks.""" - if self.verify_ssl and not self._verify_ssl_cert(url): + if self.verify_ssl and not verify_ssl_cert(url): raise ValueError(f"SSL certificate verification failed for {url}") self._sync_wait_for_rate_limit() return True diff --git a/backend/open_webui/retrieval/web/yandex.py b/backend/open_webui/retrieval/web/yandex.py index def134d996..fd1bc7274a 100644 --- a/backend/open_webui/retrieval/web/yandex.py +++ b/backend/open_webui/retrieval/web/yandex.py @@ -11,6 +11,7 @@ from fastapi import Request from open_webui.retrieval.web.main import SearchResult, get_filtered_results from open_webui.utils.headers import include_user_info_headers +from open_webui.env import FORWARD_SESSION_INFO_HEADER_CHAT_ID from xml.etree import ElementTree as ET from xml.etree.ElementTree import Element @@ -50,7 +51,7 @@ def search_yandex( chat_id = getattr(request.state, "chat_id", None) if chat_id: - headers["X-OpenWebUI-Chat-Id"] = str(chat_id) + headers[FORWARD_SESSION_INFO_HEADER_CHAT_ID] = str(chat_id) payload = {} if yandex_search_config == "" else json.loads(yandex_search_config) diff --git a/backend/open_webui/routers/channels.py b/backend/open_webui/routers/channels.py index 2c94ddd91f..4713e369ad 100644 --- a/backend/open_webui/routers/channels.py +++ b/backend/open_webui/routers/channels.py @@ -36,6 +36,7 @@ from open_webui.models.channels import ( ChannelWebhookModel, ChannelWebhookForm, ) +from open_webui.models.access_grants import AccessGrants, has_public_read_access_grant from open_webui.models.messages import ( Messages, MessageModel, @@ -60,12 +61,7 @@ from open_webui.utils.chat import generate_chat_completion from open_webui.utils.auth import get_admin_user, get_verified_user -from open_webui.utils.access_control import ( - has_access, - get_users_with_access, - get_permitted_group_and_user_ids, - has_permission, -) +from open_webui.utils.access_control import has_permission from open_webui.utils.webhook import post_webhook from open_webui.utils.channels import extract_mentions, replace_mentions from open_webui.internal.db import get_session @@ -76,6 +72,66 @@ log = logging.getLogger(__name__) router = APIRouter() +def channel_has_access( + user_id: str, + channel: ChannelModel, + permission: str = "read", + strict: bool = True, + db: Optional[Session] = None, +) -> bool: + if AccessGrants.has_access( + user_id=user_id, + resource_type="channel", + resource_id=channel.id, + permission=permission, + db=db, + ): + return True + + if ( + not strict + and permission == "write" + and has_public_read_access_grant(channel.access_grants) + ): + return True + + return False + + +def get_channel_users_with_access( + channel: ChannelModel, permission: str = "read", db: Optional[Session] = None +): + return AccessGrants.get_users_with_access( + resource_type="channel", + resource_id=channel.id, + permission=permission, + db=db, + ) + + +def get_channel_permitted_group_and_user_ids( + channel: ChannelModel, permission: str = "read" +) -> Optional[dict[str, list[str]]]: + if permission == "read" and has_public_read_access_grant(channel.access_grants): + return None + + user_ids = [] + group_ids = [] + + for grant in channel.access_grants: + if grant.permission != permission: + continue + if grant.principal_type == "group": + group_ids.append(grant.principal_id) + elif grant.principal_type == "user" and grant.principal_id != "*": + user_ids.append(grant.principal_id) + + return { + "user_ids": list(dict.fromkeys(user_ids)), + "group_ids": list(dict.fromkeys(group_ids)), + } + + ############################ # Channels Enabled Dependency ############################ @@ -418,22 +474,22 @@ async def get_channel_by_id( } ) else: - if user.role != "admin" and not has_access( - user.id, type="read", access_control=channel.access_control, db=db + if user.role != "admin" and not channel_has_access( + user.id, channel, permission="read", db=db ): raise HTTPException( status_code=status.HTTP_403_FORBIDDEN, detail=ERROR_MESSAGES.DEFAULT() ) - write_access = has_access( + write_access = channel_has_access( user.id, - type="write", - access_control=channel.access_control, + channel, + permission="write", strict=False, db=db, ) - user_count = len(get_users_with_access("read", channel.access_control)) + user_count = len(get_channel_users_with_access(channel, "read", db=db)) channel_member = Channels.get_member_by_channel_and_user_id( channel.id, user.id, db=db @@ -527,8 +583,8 @@ async def get_channel_members_by_id( filter["channel_id"] = channel.id else: filter["roles"] = ["!pending"] - permitted_ids = get_permitted_group_and_user_ids( - "read", channel.access_control + permitted_ids = get_channel_permitted_group_and_user_ids( + channel, permission="read" ) if permitted_ids: filter["user_ids"] = permitted_ids.get("user_ids") @@ -811,8 +867,8 @@ async def get_channel_messages( status_code=status.HTTP_403_FORBIDDEN, detail=ERROR_MESSAGES.DEFAULT() ) else: - if user.role != "admin" and not has_access( - user.id, type="read", access_control=channel.access_control, db=db + if user.role != "admin" and not channel_has_access( + user.id, channel, permission="read", db=db ): raise HTTPException( status_code=status.HTTP_403_FORBIDDEN, detail=ERROR_MESSAGES.DEFAULT() @@ -888,8 +944,8 @@ async def get_pinned_channel_messages( status_code=status.HTTP_403_FORBIDDEN, detail=ERROR_MESSAGES.DEFAULT() ) else: - if user.role != "admin" and not has_access( - user.id, type="read", access_control=channel.access_control, db=db + if user.role != "admin" and not channel_has_access( + user.id, channel, permission="read", db=db ): raise HTTPException( status_code=status.HTTP_403_FORBIDDEN, detail=ERROR_MESSAGES.DEFAULT() @@ -946,7 +1002,7 @@ async def get_pinned_channel_messages( async def send_notification( name, webui_url, channel, message, active_user_ids, db=None ): - users = get_users_with_access("read", channel.access_control) + users = get_channel_users_with_access(channel, "read", db=db) for user in users: if (user.id not in active_user_ids) and Channels.is_user_channel_member( @@ -1173,10 +1229,10 @@ async def new_message_handler( status_code=status.HTTP_403_FORBIDDEN, detail=ERROR_MESSAGES.DEFAULT() ) else: - if user.role != "admin" and not has_access( + if user.role != "admin" and not channel_has_access( user.id, - type="write", - access_control=channel.access_control, + channel, + permission="write", strict=False, db=db, ): @@ -1318,8 +1374,8 @@ async def get_channel_message( status_code=status.HTTP_403_FORBIDDEN, detail=ERROR_MESSAGES.DEFAULT() ) else: - if user.role != "admin" and not has_access( - user.id, type="read", access_control=channel.access_control, db=db + if user.role != "admin" and not channel_has_access( + user.id, channel, permission="read", db=db ): raise HTTPException( status_code=status.HTTP_403_FORBIDDEN, detail=ERROR_MESSAGES.DEFAULT() @@ -1372,8 +1428,8 @@ async def get_channel_message_data( status_code=status.HTTP_403_FORBIDDEN, detail=ERROR_MESSAGES.DEFAULT() ) else: - if user.role != "admin" and not has_access( - user.id, type="read", access_control=channel.access_control, db=db + if user.role != "admin" and not channel_has_access( + user.id, channel, permission="read", db=db ): raise HTTPException( status_code=status.HTTP_403_FORBIDDEN, detail=ERROR_MESSAGES.DEFAULT() @@ -1426,8 +1482,8 @@ async def pin_channel_message( status_code=status.HTTP_403_FORBIDDEN, detail=ERROR_MESSAGES.DEFAULT() ) else: - if user.role != "admin" and not has_access( - user.id, type="read", access_control=channel.access_control, db=db + if user.role != "admin" and not channel_has_access( + user.id, channel, permission="read", db=db ): raise HTTPException( status_code=status.HTTP_403_FORBIDDEN, detail=ERROR_MESSAGES.DEFAULT() @@ -1492,8 +1548,8 @@ async def get_channel_thread_messages( status_code=status.HTTP_403_FORBIDDEN, detail=ERROR_MESSAGES.DEFAULT() ) else: - if user.role != "admin" and not has_access( - user.id, type="read", access_control=channel.access_control, db=db + if user.role != "admin" and not channel_has_access( + user.id, channel, permission="read", db=db ): raise HTTPException( status_code=status.HTTP_403_FORBIDDEN, detail=ERROR_MESSAGES.DEFAULT() @@ -1577,8 +1633,8 @@ async def update_message_by_id( if ( user.role != "admin" and message.user_id != user.id - and not has_access( - user.id, type="read", access_control=channel.access_control, db=db + and not channel_has_access( + user.id, channel, permission="read", db=db ) ): raise HTTPException( @@ -1644,10 +1700,10 @@ async def add_reaction_to_message( status_code=status.HTTP_403_FORBIDDEN, detail=ERROR_MESSAGES.DEFAULT() ) else: - if user.role != "admin" and not has_access( + if user.role != "admin" and not channel_has_access( user.id, - type="write", - access_control=channel.access_control, + channel, + permission="write", strict=False, db=db, ): @@ -1723,10 +1779,10 @@ async def remove_reaction_by_id_and_user_id_and_name( status_code=status.HTTP_403_FORBIDDEN, detail=ERROR_MESSAGES.DEFAULT() ) else: - if user.role != "admin" and not has_access( + if user.role != "admin" and not channel_has_access( user.id, - type="write", - access_control=channel.access_control, + channel, + permission="write", strict=False, db=db, ): @@ -1818,10 +1874,10 @@ async def delete_message_by_id( if ( user.role != "admin" and message.user_id != user.id - and not has_access( + and not channel_has_access( user.id, - type="write", - access_control=channel.access_control, + channel, + permission="write", strict=False, db=db, ) diff --git a/backend/open_webui/routers/files.py b/backend/open_webui/routers/files.py index 8b17dc406b..4a12db5cd9 100644 --- a/backend/open_webui/routers/files.py +++ b/backend/open_webui/routers/files.py @@ -38,6 +38,7 @@ from open_webui.models.files import ( from open_webui.models.chats import Chats from open_webui.models.knowledge import Knowledges from open_webui.models.groups import Groups +from open_webui.models.access_grants import AccessGrants from open_webui.routers.retrieval import ProcessFileForm, process_file @@ -47,7 +48,6 @@ from open_webui.storage.provider import Storage from open_webui.utils.auth import get_admin_user, get_verified_user -from open_webui.utils.access_control import has_access from open_webui.utils.misc import strict_match_mime_type from pydantic import BaseModel @@ -82,8 +82,13 @@ def has_access_to_file( group.id for group in Groups.get_groups_by_member_id(user.id, db=db) } for knowledge_base in knowledge_bases: - if knowledge_base.user_id == user.id or has_access( - user.id, access_type, knowledge_base.access_control, user_group_ids, db=db + if knowledge_base.user_id == user.id or AccessGrants.has_access( + user_id=user.id, + resource_type="knowledge", + resource_id=knowledge_base.id, + permission=access_type, + user_group_ids=user_group_ids, + db=db, ): return True diff --git a/backend/open_webui/routers/groups.py b/backend/open_webui/routers/groups.py index cc0cb8f5a3..a46e05473f 100755 --- a/backend/open_webui/routers/groups.py +++ b/backend/open_webui/routers/groups.py @@ -7,6 +7,7 @@ from open_webui.models.users import Users, UserInfoResponse from open_webui.models.groups import ( Groups, GroupForm, + GroupInfoResponse, GroupUpdateForm, GroupResponse, UserIdsForm, @@ -104,6 +105,23 @@ async def get_group_by_id( ) +@router.get("/id/{id}/info", response_model=Optional[GroupInfoResponse]) +async def get_group_info_by_id( + id: str, user=Depends(get_verified_user), db: Session = Depends(get_session) +): + group = Groups.get_group_by_id(id, db=db) + if group: + return GroupInfoResponse( + **group.model_dump(), + member_count=Groups.get_group_member_count_by_id(group.id, db=db), + ) + else: + raise HTTPException( + status_code=status.HTTP_401_UNAUTHORIZED, + detail=ERROR_MESSAGES.NOT_FOUND, + ) + + ############################ # ExportGroupById ############################ diff --git a/backend/open_webui/routers/knowledge.py b/backend/open_webui/routers/knowledge.py index 96136d6898..7fff5e5f2a 100644 --- a/backend/open_webui/routers/knowledge.py +++ b/backend/open_webui/routers/knowledge.py @@ -29,7 +29,8 @@ from open_webui.storage.provider import Storage from open_webui.constants import ERROR_MESSAGES from open_webui.utils.auth import get_verified_user, get_admin_user -from open_webui.utils.access_control import has_access, has_permission +from open_webui.utils.access_control import has_permission +from open_webui.models.access_grants import AccessGrants, has_public_read_access_grant from open_webui.config import BYPASS_ADMIN_ACCESS_CONTROL @@ -133,8 +134,12 @@ async def get_knowledge_bases( write_access=( user.id == knowledge_base.user_id or (user.role == "admin" and BYPASS_ADMIN_ACCESS_CONTROL) - or has_access( - user.id, "write", knowledge_base.access_control, db=db + or AccessGrants.has_access( + user_id=user.id, + resource_type="knowledge", + resource_id=knowledge_base.id, + permission="write", + db=db, ) ), ) @@ -180,8 +185,12 @@ async def search_knowledge_bases( write_access=( user.id == knowledge_base.user_id or (user.role == "admin" and BYPASS_ADMIN_ACCESS_CONTROL) - or has_access( - user.id, "write", knowledge_base.access_control, db=db + or AccessGrants.has_access( + user_id=user.id, + resource_type="knowledge", + resource_id=knowledge_base.id, + permission="write", + db=db, ) ), ) @@ -243,14 +252,14 @@ async def create_new_knowledge( # Check if user can share publicly if ( user.role != "admin" - and form_data.access_control == None + and has_public_read_access_grant(form_data.access_grants) and not has_permission( user.id, "sharing.public_knowledge", request.app.state.config.USER_PERMISSIONS, ) ): - form_data.access_control = {} + form_data.access_grants = [] knowledge = Knowledges.insert_new_knowledge(user.id, form_data) @@ -387,7 +396,13 @@ async def get_knowledge_by_id( if ( user.role == "admin" or knowledge.user_id == user.id - or has_access(user.id, "read", knowledge.access_control, db=db) + or AccessGrants.has_access( + user_id=user.id, + resource_type="knowledge", + resource_id=knowledge.id, + permission="read", + db=db, + ) ): return KnowledgeFilesResponse( @@ -395,7 +410,13 @@ async def get_knowledge_by_id( write_access=( user.id == knowledge.user_id or (user.role == "admin" and BYPASS_ADMIN_ACCESS_CONTROL) - or has_access(user.id, "write", knowledge.access_control, db=db) + or AccessGrants.has_access( + user_id=user.id, + resource_type="knowledge", + resource_id=knowledge.id, + permission="write", + db=db, + ) ), ) else: @@ -435,7 +456,12 @@ async def update_knowledge_by_id( # Is the user the original creator, in a group with write access, or an admin if ( knowledge.user_id != user.id - and not has_access(user.id, "write", knowledge.access_control) + and not AccessGrants.has_access( + user_id=user.id, + resource_type="knowledge", + resource_id=knowledge.id, + permission="write", + ) and user.role != "admin" ): raise HTTPException( @@ -446,14 +472,14 @@ async def update_knowledge_by_id( # Check if user can share publicly if ( user.role != "admin" - and form_data.access_control == None + and has_public_read_access_grant(form_data.access_grants) and not has_permission( user.id, "sharing.public_knowledge", request.app.state.config.USER_PERMISSIONS, ) ): - form_data.access_control = {} + form_data.access_grants = [] knowledge = Knowledges.update_knowledge_by_id(id=id, form_data=form_data) if knowledge: @@ -502,7 +528,13 @@ async def get_knowledge_files_by_id( if not ( user.role == "admin" or knowledge.user_id == user.id - or has_access(user.id, "read", knowledge.access_control, db=db) + or AccessGrants.has_access( + user_id=user.id, + resource_type="knowledge", + resource_id=knowledge.id, + permission="read", + db=db, + ) ): raise HTTPException( status_code=status.HTTP_400_BAD_REQUEST, @@ -555,7 +587,13 @@ def add_file_to_knowledge_by_id( if ( knowledge.user_id != user.id - and not has_access(user.id, "write", knowledge.access_control, db=db) + and not AccessGrants.has_access( + user_id=user.id, + resource_type="knowledge", + resource_id=knowledge.id, + permission="write", + db=db, + ) and user.role != "admin" ): raise HTTPException( @@ -624,7 +662,13 @@ def update_file_from_knowledge_by_id( if ( knowledge.user_id != user.id - and not has_access(user.id, "write", knowledge.access_control, db=db) + and not AccessGrants.has_access( + user_id=user.id, + resource_type="knowledge", + resource_id=knowledge.id, + permission="write", + db=db, + ) and user.role != "admin" ): @@ -693,7 +737,13 @@ def remove_file_from_knowledge_by_id( if ( knowledge.user_id != user.id - and not has_access(user.id, "write", knowledge.access_control, db=db) + and not AccessGrants.has_access( + user_id=user.id, + resource_type="knowledge", + resource_id=knowledge.id, + permission="write", + db=db, + ) and user.role != "admin" ): raise HTTPException( @@ -770,7 +820,13 @@ async def delete_knowledge_by_id( if ( knowledge.user_id != user.id - and not has_access(user.id, "write", knowledge.access_control, db=db) + and not AccessGrants.has_access( + user_id=user.id, + resource_type="knowledge", + resource_id=knowledge.id, + permission="write", + db=db, + ) and user.role != "admin" ): raise HTTPException( @@ -802,7 +858,7 @@ async def delete_knowledge_by_id( base_model_id=model.base_model_id, meta=model.meta, params=model.params, - access_control=model.access_control, + access_grants=model.access_grants, is_active=model.is_active, ) Models.update_model_by_id(model.id, model_form, db=db) @@ -839,7 +895,13 @@ async def reset_knowledge_by_id( if ( knowledge.user_id != user.id - and not has_access(user.id, "write", knowledge.access_control, db=db) + and not AccessGrants.has_access( + user_id=user.id, + resource_type="knowledge", + resource_id=knowledge.id, + permission="write", + db=db, + ) and user.role != "admin" ): raise HTTPException( @@ -882,7 +944,13 @@ async def add_files_to_knowledge_batch( if ( knowledge.user_id != user.id - and not has_access(user.id, "write", knowledge.access_control, db=db) + and not AccessGrants.has_access( + user_id=user.id, + resource_type="knowledge", + resource_id=knowledge.id, + permission="write", + db=db, + ) and user.role != "admin" ): raise HTTPException( @@ -890,17 +958,19 @@ async def add_files_to_knowledge_batch( detail=ERROR_MESSAGES.ACCESS_PROHIBITED, ) - # Get files content + # Batch-fetch all files to avoid N+1 queries log.info(f"files/batch/add - {len(form_data)} files") - files: List[FileModel] = [] - for form in form_data: - file = Files.get_file_by_id(form.file_id, db=db) - if not file: - raise HTTPException( - status_code=status.HTTP_400_BAD_REQUEST, - detail=f"File {form.file_id} not found", - ) - files.append(file) + file_ids = [form.file_id for form in form_data] + files = Files.get_files_by_ids(file_ids, db=db) + + # Verify all requested files were found + found_ids = {file.id for file in files} + missing_ids = [fid for fid in file_ids if fid not in found_ids] + if missing_ids: + raise HTTPException( + status_code=status.HTTP_400_BAD_REQUEST, + detail=f"File {missing_ids[0]} not found", + ) # Process files try: diff --git a/backend/open_webui/routers/models.py b/backend/open_webui/routers/models.py index 85f0fb4f64..b2bccc1958 100644 --- a/backend/open_webui/routers/models.py +++ b/backend/open_webui/routers/models.py @@ -15,6 +15,7 @@ from open_webui.models.models import ( ModelAccessResponse, Models, ) +from open_webui.models.access_grants import AccessGrants from pydantic import BaseModel from open_webui.constants import ERROR_MESSAGES @@ -30,7 +31,7 @@ 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_access, has_permission +from open_webui.utils.access_control import has_permission from open_webui.config import BYPASS_ADMIN_ACCESS_CONTROL, STATIC_DIR from open_webui.internal.db import get_session from sqlalchemy.orm import Session @@ -98,7 +99,13 @@ async def get_models( write_access=( (user.role == "admin" and BYPASS_ADMIN_ACCESS_CONTROL) or user.id == model.user_id - or has_access(user.id, "write", model.access_control, db=db) + or AccessGrants.has_access( + user_id=user.id, + resource_type="model", + resource_id=model.id, + permission="write", + db=db, + ) ), ) for model in result.items @@ -315,14 +322,26 @@ async def get_model_by_id( if ( (user.role == "admin" and BYPASS_ADMIN_ACCESS_CONTROL) or model.user_id == user.id - or has_access(user.id, "read", model.access_control, db=db) + or AccessGrants.has_access( + user_id=user.id, + resource_type="model", + resource_id=model.id, + permission="read", + db=db, + ) ): return ModelAccessResponse( **model.model_dump(), write_access=( (user.role == "admin" and BYPASS_ADMIN_ACCESS_CONTROL) or user.id == model.user_id - or has_access(user.id, "write", model.access_control, db=db) + or AccessGrants.has_access( + user_id=user.id, + resource_type="model", + resource_id=model.id, + permission="write", + db=db, + ) ), ) else: @@ -393,7 +412,13 @@ async def toggle_model_by_id( if ( user.role == "admin" or model.user_id == user.id - or has_access(user.id, "write", model.access_control, db=db) + or AccessGrants.has_access( + user_id=user.id, + resource_type="model", + resource_id=model.id, + permission="write", + db=db, + ) ): model = Models.toggle_model_by_id(id, db=db) @@ -436,7 +461,13 @@ async def update_model_by_id( if ( model.user_id != user.id - and not has_access(user.id, "write", model.access_control, db=db) + and not AccessGrants.has_access( + user_id=user.id, + resource_type="model", + resource_id=model.id, + permission="write", + db=db, + ) and user.role != "admin" ): raise HTTPException( @@ -471,7 +502,13 @@ async def delete_model_by_id( if ( user.role != "admin" and model.user_id != user.id - and not has_access(user.id, "write", model.access_control, db=db) + and not AccessGrants.has_access( + user_id=user.id, + resource_type="model", + resource_id=model.id, + permission="write", + db=db, + ) ): raise HTTPException( status_code=status.HTTP_401_UNAUTHORIZED, diff --git a/backend/open_webui/routers/notes.py b/backend/open_webui/routers/notes.py index 56730e2b6a..321b06fcda 100644 --- a/backend/open_webui/routers/notes.py +++ b/backend/open_webui/routers/notes.py @@ -27,7 +27,8 @@ from open_webui.constants import ERROR_MESSAGES from open_webui.utils.auth import get_admin_user, get_verified_user -from open_webui.utils.access_control import has_access, has_permission +from open_webui.utils.access_control import has_permission +from open_webui.models.access_grants import AccessGrants, has_public_read_access_grant from open_webui.internal.db import get_session from sqlalchemy.orm import Session @@ -200,8 +201,12 @@ async def get_note_by_id( if user.role != "admin" and ( user.id != note.user_id and ( - not has_access( - user.id, type="read", access_control=note.access_control, db=db + not AccessGrants.has_access( + user_id=user.id, + resource_type="note", + resource_id=note.id, + permission="read", + db=db, ) ) ): @@ -212,13 +217,14 @@ async def get_note_by_id( write_access = ( user.role == "admin" or (user.id == note.user_id) - or has_access( - user.id, - type="write", - access_control=note.access_control, - strict=False, + or AccessGrants.has_access( + user_id=user.id, + resource_type="note", + resource_id=note.id, + permission="write", db=db, ) + or has_public_read_access_grant(note.access_grants) ) return NoteResponse(**note.model_dump(), write_access=write_access) @@ -253,8 +259,12 @@ async def update_note_by_id( if user.role != "admin" and ( user.id != note.user_id - and not has_access( - user.id, type="write", access_control=note.access_control, db=db + and not AccessGrants.has_access( + user_id=user.id, + resource_type="note", + resource_id=note.id, + permission="write", + db=db, ) ): raise HTTPException( @@ -264,7 +274,7 @@ async def update_note_by_id( # Check if user can share publicly if ( user.role != "admin" - and form_data.access_control == None + and has_public_read_access_grant(form_data.access_grants) and not has_permission( user.id, "sharing.public_notes", @@ -272,7 +282,7 @@ async def update_note_by_id( db=db, ) ): - form_data.access_control = {} + form_data.access_grants = [] try: note = Notes.update_note_by_id(id, form_data, db=db) @@ -318,8 +328,12 @@ async def delete_note_by_id( if user.role != "admin" and ( user.id != note.user_id - and not has_access( - user.id, type="write", access_control=note.access_control, db=db + and not AccessGrants.has_access( + user_id=user.id, + resource_type="note", + resource_id=note.id, + permission="write", + db=db, ) ): raise HTTPException( diff --git a/backend/open_webui/routers/ollama.py b/backend/open_webui/routers/ollama.py index cfbdbcac08..4f43c41d30 100644 --- a/backend/open_webui/routers/ollama.py +++ b/backend/open_webui/routers/ollama.py @@ -23,6 +23,7 @@ from open_webui.models.users import UserModel from open_webui.env import ( ENABLE_FORWARD_USER_INFO_HEADERS, + FORWARD_SESSION_INFO_HEADER_CHAT_ID, ) from fastapi import ( @@ -44,6 +45,7 @@ from open_webui.internal.db import get_session from open_webui.models.models import Models +from open_webui.models.access_grants import AccessGrants from open_webui.utils.misc import ( calculate_sha256, ) @@ -53,9 +55,6 @@ from open_webui.utils.payload import ( apply_system_prompt_to_body, ) from open_webui.utils.auth import get_admin_user, get_verified_user -from open_webui.utils.access_control import has_access - - from open_webui.config import ( UPLOAD_DIR, ) @@ -137,7 +136,7 @@ async def send_post_request( if ENABLE_FORWARD_USER_INFO_HEADERS and user: headers = include_user_info_headers(headers, user) if metadata and metadata.get("chat_id"): - headers["X-OpenWebUI-Chat-Id"] = metadata.get("chat_id") + headers[FORWARD_SESSION_INFO_HEADER_CHAT_ID] = metadata.get("chat_id") r = await session.post( url, @@ -430,8 +429,12 @@ async def get_filtered_models(models, user, db=None): for model in models.get("models", []): model_info = Models.get_model_by_id(model["model"], db=db) if model_info: - if user.id == model_info.user_id or has_access( - user.id, type="read", access_control=model_info.access_control, db=db + if user.id == model_info.user_id or AccessGrants.has_access( + user_id=user.id, + resource_type="model", + resource_id=model_info.id, + permission="read", + db=db, ): filtered_models.append(model) return filtered_models @@ -442,6 +445,9 @@ async def get_filtered_models(models, user, db=None): async def get_ollama_tags( request: Request, url_idx: Optional[int] = None, user=Depends(get_verified_user) ): + if not request.app.state.config.ENABLE_OLLAMA_API: + raise HTTPException(status_code=503, detail="Ollama API is disabled") + models = [] if url_idx is None: @@ -706,6 +712,9 @@ async def pull_model( url_idx: int = 0, user=Depends(get_admin_user), ): + if not request.app.state.config.ENABLE_OLLAMA_API: + raise HTTPException(status_code=503, detail="Ollama API is disabled") + form_data = form_data.model_dump(exclude_none=True) form_data["model"] = form_data.get("model", form_data.get("name")) @@ -737,6 +746,9 @@ async def push_model( url_idx: Optional[int] = None, user=Depends(get_admin_user), ): + if not request.app.state.config.ENABLE_OLLAMA_API: + raise HTTPException(status_code=503, detail="Ollama API is disabled") + if url_idx is None: await get_all_models(request, user=user) models = request.app.state.OLLAMA_MODELS @@ -776,6 +788,9 @@ async def create_model( url_idx: int = 0, user=Depends(get_admin_user), ): + if not request.app.state.config.ENABLE_OLLAMA_API: + raise HTTPException(status_code=503, detail="Ollama API is disabled") + log.debug(f"form_data: {form_data}") url = request.app.state.config.OLLAMA_BASE_URLS[url_idx] @@ -800,6 +815,9 @@ async def copy_model( url_idx: Optional[int] = None, user=Depends(get_admin_user), ): + if not request.app.state.config.ENABLE_OLLAMA_API: + raise HTTPException(status_code=503, detail="Ollama API is disabled") + if url_idx is None: await get_all_models(request, user=user) models = request.app.state.OLLAMA_MODELS @@ -860,6 +878,9 @@ async def delete_model( url_idx: Optional[int] = None, user=Depends(get_admin_user), ): + if not request.app.state.config.ENABLE_OLLAMA_API: + raise HTTPException(status_code=503, detail="Ollama API is disabled") + form_data = form_data.model_dump(exclude_none=True) form_data["model"] = form_data.get("model", form_data.get("name")) @@ -922,6 +943,9 @@ async def delete_model( async def show_model_info( request: Request, form_data: ModelNameForm, user=Depends(get_verified_user) ): + if not request.app.state.config.ENABLE_OLLAMA_API: + raise HTTPException(status_code=503, detail="Ollama API is disabled") + form_data = form_data.model_dump(exclude_none=True) form_data["model"] = form_data.get("model", form_data.get("name")) @@ -994,6 +1018,9 @@ async def embed( url_idx: Optional[int] = None, user=Depends(get_verified_user), ): + if not request.app.state.config.ENABLE_OLLAMA_API: + raise HTTPException(status_code=503, detail="Ollama API is disabled") + log.info(f"generate_ollama_batch_embeddings {form_data}") if url_idx is None: @@ -1076,6 +1103,9 @@ async def embeddings( url_idx: Optional[int] = None, user=Depends(get_verified_user), ): + if not request.app.state.config.ENABLE_OLLAMA_API: + raise HTTPException(status_code=503, detail="Ollama API is disabled") + log.info(f"generate_ollama_embeddings {form_data}") if url_idx is None: @@ -1166,6 +1196,9 @@ async def generate_completion( url_idx: Optional[int] = None, user=Depends(get_verified_user), ): + if not request.app.state.config.ENABLE_OLLAMA_API: + raise HTTPException(status_code=503, detail="Ollama API is disabled") + if url_idx is None: await get_all_models(request, user=user) models = request.app.state.OLLAMA_MODELS @@ -1258,8 +1291,11 @@ async def generate_chat_completion( bypass_filter: Optional[bool] = False, bypass_system_prompt: bool = False, ): + if not request.app.state.config.ENABLE_OLLAMA_API: + raise HTTPException(status_code=503, detail="Ollama API is disabled") + # NOTE: We intentionally do NOT use Depends(get_session) here. - # Database operations (get_model_by_id, has_access) manage their own short-lived sessions. + # Database operations (get_model_by_id, AccessGrants.has_access) manage their own short-lived sessions. # This prevents holding a connection during the entire LLM call (30-60+ seconds), # which would exhaust the connection pool under concurrent load. if BYPASS_MODEL_ACCESS_CONTROL: @@ -1306,10 +1342,11 @@ async def generate_chat_completion( if not bypass_filter and user.role == "user": if not ( user.id == model_info.user_id - or has_access( - user.id, - type="read", - access_control=model_info.access_control, + or AccessGrants.has_access( + user_id=user.id, + resource_type="model", + resource_id=model_info.id, + permission="read", ) ): raise HTTPException( @@ -1383,7 +1420,7 @@ async def generate_openai_completion( user=Depends(get_verified_user), ): # NOTE: We intentionally do NOT use Depends(get_session) here. - # Database operations (get_model_by_id, has_access) manage their own short-lived sessions. + # Database operations (get_model_by_id, AccessGrants.has_access) manage their own short-lived sessions. # This prevents holding a connection during the entire LLM call (30-60+ seconds), # which would exhaust the connection pool under concurrent load. metadata = form_data.pop("metadata", None) @@ -1418,10 +1455,11 @@ async def generate_openai_completion( if user.role == "user": if not ( user.id == model_info.user_id - or has_access( - user.id, - type="read", - access_control=model_info.access_control, + or AccessGrants.has_access( + user_id=user.id, + resource_type="model", + resource_id=model_info.id, + permission="read", ) ): raise HTTPException( @@ -1468,7 +1506,7 @@ async def generate_openai_chat_completion( user=Depends(get_verified_user), ): # NOTE: We intentionally do NOT use Depends(get_session) here. - # Database operations (get_model_by_id, has_access) manage their own short-lived sessions. + # Database operations (get_model_by_id, AccessGrants.has_access) manage their own short-lived sessions. # This prevents holding a connection during the entire LLM call (30-60+ seconds), # which would exhaust the connection pool under concurrent load. metadata = form_data.pop("metadata", None) @@ -1507,10 +1545,11 @@ async def generate_openai_chat_completion( if user.role == "user": if not ( user.id == model_info.user_id - or has_access( - user.id, - type="read", - access_control=model_info.access_control, + or AccessGrants.has_access( + user_id=user.id, + resource_type="model", + resource_id=model_info.id, + permission="read", ) ): raise HTTPException( @@ -1608,10 +1647,11 @@ async def get_openai_models( for model in models: model_info = Models.get_model_by_id(model["id"], db=db) if model_info: - if user.id == model_info.user_id or has_access( - user.id, - type="read", - access_control=model_info.access_control, + if user.id == model_info.user_id or AccessGrants.has_access( + user_id=user.id, + resource_type="model", + resource_id=model_info.id, + permission="read", db=db, ): filtered_models.append(model) diff --git a/backend/open_webui/routers/openai.py b/backend/open_webui/routers/openai.py index b1d31afb8f..d8ab50221f 100644 --- a/backend/open_webui/routers/openai.py +++ b/backend/open_webui/routers/openai.py @@ -24,6 +24,7 @@ from sqlalchemy.orm import Session from open_webui.internal.db import get_session from open_webui.models.models import Models +from open_webui.models.access_grants import AccessGrants from open_webui.config import ( CACHE_DIR, ) @@ -33,6 +34,7 @@ from open_webui.env import ( AIOHTTP_CLIENT_TIMEOUT, AIOHTTP_CLIENT_TIMEOUT_MODEL_LIST, ENABLE_FORWARD_USER_INFO_HEADERS, + FORWARD_SESSION_INFO_HEADER_CHAT_ID, BYPASS_MODEL_ACCESS_CONTROL, ) from open_webui.models.users import UserModel @@ -50,7 +52,6 @@ from open_webui.utils.misc import ( ) from open_webui.utils.auth import get_admin_user, get_verified_user -from open_webui.utils.access_control import has_access from open_webui.utils.headers import include_user_info_headers @@ -142,7 +143,7 @@ async def get_headers_and_cookies( if ENABLE_FORWARD_USER_INFO_HEADERS and user: headers = include_user_info_headers(headers, user) if metadata and metadata.get("chat_id"): - headers["X-OpenWebUI-Chat-Id"] = metadata.get("chat_id") + headers[FORWARD_SESSION_INFO_HEADER_CHAT_ID] = metadata.get("chat_id") token = None auth_type = config.get("auth_type") @@ -462,8 +463,12 @@ async def get_filtered_models(models, user, db=None): for model in models.get("data", []): model_info = Models.get_model_by_id(model["id"], db=db) if model_info: - if user.id == model_info.user_id or has_access( - user.id, type="read", access_control=model_info.access_control, db=db + if user.id == model_info.user_id or AccessGrants.has_access( + user_id=user.id, + resource_type="model", + resource_id=model_info.id, + permission="read", + db=db, ): filtered_models.append(model) return filtered_models @@ -544,6 +549,9 @@ async def get_all_models(request: Request, user: UserModel) -> dict[str, list]: async def get_models( request: Request, url_idx: Optional[int] = None, user=Depends(get_verified_user) ): + if not request.app.state.config.ENABLE_OPENAI_API: + raise HTTPException(status_code=503, detail="OpenAI API is disabled") + models = { "data": [], } @@ -854,6 +862,33 @@ def convert_to_responses_payload(payload: dict) -> dict: 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"): + 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": ...} + if "tools" in responses_payload and isinstance(responses_payload["tools"], list): + converted_tools = [] + for tool in responses_payload["tools"]: + if isinstance(tool, dict) and "function" in tool: + func = tool["function"] + converted_tool = {"type": tool.get("type", "function")} + if isinstance(func, dict): + converted_tool["name"] = func.get("name", "") + if "description" in func: + converted_tool["description"] = func["description"] + if "parameters" in func: + converted_tool["parameters"] = func["parameters"] + if "strict" in func: + converted_tool["strict"] = func["strict"] + converted_tools.append(converted_tool) + else: + # Already in correct format or unknown format, pass through + converted_tools.append(tool) + responses_payload["tools"] = converted_tools + return responses_payload @@ -876,7 +911,7 @@ async def generate_chat_completion( bypass_system_prompt: bool = False, ): # NOTE: We intentionally do NOT use Depends(get_session) here. - # Database operations (get_model_by_id, has_access) manage their own short-lived sessions. + # Database operations (get_model_by_id, AccessGrants.has_access) manage their own short-lived sessions. # This prevents holding a connection during the entire LLM call (30-60+ seconds), # which would exhaust the connection pool under concurrent load. if BYPASS_MODEL_ACCESS_CONTROL: @@ -914,10 +949,11 @@ async def generate_chat_completion( if not bypass_filter and user.role == "user": if not ( user.id == model_info.user_id - or has_access( - user.id, - type="read", - access_control=model_info.access_control, + or AccessGrants.has_access( + user_id=user.id, + resource_type="model", + resource_id=model_info.id, + permission="read", ) ): raise HTTPException( diff --git a/backend/open_webui/routers/prompts.py b/backend/open_webui/routers/prompts.py index fc24ccaf4b..77c2e84fc4 100644 --- a/backend/open_webui/routers/prompts.py +++ b/backend/open_webui/routers/prompts.py @@ -9,6 +9,7 @@ from open_webui.models.prompts import ( PromptModel, Prompts, ) +from open_webui.models.access_grants import AccessGrants from open_webui.models.groups import Groups from open_webui.models.prompt_history import ( PromptHistories, @@ -17,7 +18,7 @@ from open_webui.models.prompt_history import ( ) from open_webui.constants import ERROR_MESSAGES from open_webui.utils.auth import get_admin_user, get_verified_user -from open_webui.utils.access_control import has_access, has_permission +from open_webui.utils.access_control import has_permission from open_webui.config import BYPASS_ADMIN_ACCESS_CONTROL from open_webui.internal.db import get_session from sqlalchemy.orm import Session @@ -115,7 +116,13 @@ async def get_prompt_list( write_access=( (user.role == "admin" and BYPASS_ADMIN_ACCESS_CONTROL) or user.id == prompt.user_id - or has_access(user.id, "write", prompt.access_control, db=db) + or AccessGrants.has_access( + user_id=user.id, + resource_type="prompt", + resource_id=prompt.id, + permission="write", + db=db, + ) ), ) for prompt in result.items @@ -186,14 +193,26 @@ async def get_prompt_by_command( if ( user.role == "admin" or prompt.user_id == user.id - or has_access(user.id, "read", prompt.access_control, db=db) + or AccessGrants.has_access( + user_id=user.id, + resource_type="prompt", + resource_id=prompt.id, + permission="read", + db=db, + ) ): return PromptAccessResponse( **prompt.model_dump(), write_access=( (user.role == "admin" and BYPASS_ADMIN_ACCESS_CONTROL) or user.id == prompt.user_id - or has_access(user.id, "write", prompt.access_control, db=db) + or AccessGrants.has_access( + user_id=user.id, + resource_type="prompt", + resource_id=prompt.id, + permission="write", + db=db, + ) ), ) @@ -218,14 +237,26 @@ async def get_prompt_by_id( if ( user.role == "admin" or prompt.user_id == user.id - or has_access(user.id, "read", prompt.access_control, db=db) + or AccessGrants.has_access( + user_id=user.id, + resource_type="prompt", + resource_id=prompt.id, + permission="read", + db=db, + ) ): return PromptAccessResponse( **prompt.model_dump(), write_access=( (user.role == "admin" and BYPASS_ADMIN_ACCESS_CONTROL) or user.id == prompt.user_id - or has_access(user.id, "write", prompt.access_control, db=db) + or AccessGrants.has_access( + user_id=user.id, + resource_type="prompt", + resource_id=prompt.id, + permission="write", + db=db, + ) ), ) @@ -258,7 +289,13 @@ async def update_prompt_by_id( # Is the user the original creator, in a group with write access, or an admin if ( prompt.user_id != user.id - and not has_access(user.id, "write", prompt.access_control, db=db) + and not AccessGrants.has_access( + user_id=user.id, + resource_type="prompt", + resource_id=prompt.id, + permission="write", + db=db, + ) and user.role != "admin" ): raise HTTPException( @@ -311,7 +348,13 @@ async def update_prompt_metadata( if ( prompt.user_id != user.id - and not has_access(user.id, "write", prompt.access_control, db=db) + and not AccessGrants.has_access( + user_id=user.id, + resource_type="prompt", + resource_id=prompt.id, + permission="write", + db=db, + ) and user.role != "admin" ): raise HTTPException( @@ -356,7 +399,13 @@ async def set_prompt_version( if ( prompt.user_id != user.id - and not has_access(user.id, "write", prompt.access_control, db=db) + and not AccessGrants.has_access( + user_id=user.id, + resource_type="prompt", + resource_id=prompt.id, + permission="write", + db=db, + ) and user.role != "admin" ): raise HTTPException( @@ -395,7 +444,13 @@ async def delete_prompt_by_id( if ( prompt.user_id != user.id - and not has_access(user.id, "write", prompt.access_control, db=db) + and not AccessGrants.has_access( + user_id=user.id, + resource_type="prompt", + resource_id=prompt.id, + permission="write", + db=db, + ) and user.role != "admin" ): raise HTTPException( @@ -434,7 +489,13 @@ async def get_prompt_history( if not ( user.role == "admin" or prompt.user_id == user.id - or has_access(user.id, "read", prompt.access_control, db=db) + or AccessGrants.has_access( + user_id=user.id, + resource_type="prompt", + resource_id=prompt.id, + permission="read", + db=db, + ) ): raise HTTPException( status_code=status.HTTP_401_UNAUTHORIZED, @@ -469,7 +530,13 @@ async def get_prompt_history_entry( if not ( user.role == "admin" or prompt.user_id == user.id - or has_access(user.id, "read", prompt.access_control, db=db) + or AccessGrants.has_access( + user_id=user.id, + resource_type="prompt", + resource_id=prompt.id, + permission="read", + db=db, + ) ): raise HTTPException( status_code=status.HTTP_401_UNAUTHORIZED, @@ -508,7 +575,13 @@ async def delete_prompt_history_entry( if not ( user.role == "admin" or prompt.user_id == user.id - or has_access(user.id, "write", prompt.access_control, db=db) + or AccessGrants.has_access( + user_id=user.id, + resource_type="prompt", + resource_id=prompt.id, + permission="write", + db=db, + ) ): raise HTTPException( status_code=status.HTTP_401_UNAUTHORIZED, @@ -553,7 +626,13 @@ async def get_prompt_diff( if not ( user.role == "admin" or prompt.user_id == user.id - or has_access(user.id, "read", prompt.access_control, db=db) + or AccessGrants.has_access( + user_id=user.id, + resource_type="prompt", + resource_id=prompt.id, + permission="read", + db=db, + ) ): raise HTTPException( status_code=status.HTTP_401_UNAUTHORIZED, diff --git a/backend/open_webui/routers/retrieval.py b/backend/open_webui/routers/retrieval.py index 72bd0a42b0..01b1adac2c 100644 --- a/backend/open_webui/routers/retrieval.py +++ b/backend/open_webui/routers/retrieval.py @@ -1056,13 +1056,25 @@ async def update_rag_config( ) # File upload settings - request.app.state.config.FILE_MAX_SIZE = form_data.FILE_MAX_SIZE - request.app.state.config.FILE_MAX_COUNT = form_data.FILE_MAX_COUNT + request.app.state.config.FILE_MAX_SIZE = ( + form_data.FILE_MAX_SIZE + if form_data.FILE_MAX_SIZE is not None + else request.app.state.config.FILE_MAX_SIZE + ) + request.app.state.config.FILE_MAX_COUNT = ( + form_data.FILE_MAX_COUNT + if form_data.FILE_MAX_COUNT is not None + else request.app.state.config.FILE_MAX_COUNT + ) request.app.state.config.FILE_IMAGE_COMPRESSION_WIDTH = ( form_data.FILE_IMAGE_COMPRESSION_WIDTH + if form_data.FILE_IMAGE_COMPRESSION_WIDTH is not None + else request.app.state.config.FILE_IMAGE_COMPRESSION_WIDTH ) request.app.state.config.FILE_IMAGE_COMPRESSION_HEIGHT = ( form_data.FILE_IMAGE_COMPRESSION_HEIGHT + if form_data.FILE_IMAGE_COMPRESSION_HEIGHT is not None + else request.app.state.config.FILE_IMAGE_COMPRESSION_HEIGHT ) request.app.state.config.ALLOWED_FILE_EXTENSIONS = ( form_data.ALLOWED_FILE_EXTENSIONS diff --git a/backend/open_webui/routers/tasks.py b/backend/open_webui/routers/tasks.py index 37a80e2ce0..a89e28d3af 100644 --- a/backend/open_webui/routers/tasks.py +++ b/backend/open_webui/routers/tasks.py @@ -49,6 +49,21 @@ router = APIRouter() ################################## +class ActiveChatsForm(BaseModel): + chat_ids: list[str] + + +@router.post("/active/chats") +async def check_active_chats( + request: Request, form_data: ActiveChatsForm, user=Depends(get_verified_user) +): + """Check which chat IDs have active tasks.""" + from open_webui.tasks import get_active_chat_ids + + active = await get_active_chat_ids(request.app.state.redis, form_data.chat_ids) + return {"active_chat_ids": active} + + @router.get("/config") async def get_task_config(request: Request, user=Depends(get_verified_user)): return { diff --git a/backend/open_webui/routers/tools.py b/backend/open_webui/routers/tools.py index 7f9b23c7ce..015bde232a 100644 --- a/backend/open_webui/routers/tools.py +++ b/backend/open_webui/routers/tools.py @@ -21,6 +21,7 @@ from open_webui.models.tools import ( ToolAccessResponse, Tools, ) +from open_webui.models.access_grants import AccessGrants from open_webui.utils.plugin import ( load_tool_module_by_id, replace_imports, @@ -156,7 +157,24 @@ async def get_tools( tool for tool in tools if tool.user_id == user.id - or has_access(user.id, "read", tool.access_control, user_group_ids, db=db) + or ( + has_access( + user.id, + "read", + getattr(tool, "access_control", None), + user_group_ids, + db=db, + ) + if str(tool.id).startswith("server:") + else AccessGrants.has_access( + user_id=user.id, + resource_type="tool", + resource_id=tool.id, + permission="read", + user_group_ids=user_group_ids, + db=db, + ) + ) ] return tools @@ -181,7 +199,13 @@ async def get_tool_list( write_access=( (user.role == "admin" and BYPASS_ADMIN_ACCESS_CONTROL) or user.id == tool.user_id - or has_access(user.id, "write", tool.access_control, db=db) + or AccessGrants.has_access( + user_id=user.id, + resource_type="tool", + resource_id=tool.id, + permission="write", + db=db, + ) ), ) for tool in tools @@ -382,14 +406,26 @@ async def get_tools_by_id( if ( user.role == "admin" or tools.user_id == user.id - or has_access(user.id, "read", tools.access_control, db=db) + or AccessGrants.has_access( + user_id=user.id, + resource_type="tool", + resource_id=tools.id, + permission="read", + db=db, + ) ): return ToolAccessResponse( **tools.model_dump(), write_access=( (user.role == "admin" and BYPASS_ADMIN_ACCESS_CONTROL) or user.id == tools.user_id - or has_access(user.id, "write", tools.access_control, db=db) + or AccessGrants.has_access( + user_id=user.id, + resource_type="tool", + resource_id=tools.id, + permission="write", + db=db, + ) ), ) else: @@ -427,7 +463,13 @@ async def update_tools_by_id( # Is the user the original creator, in a group with write access, or an admin if ( tools.user_id != user.id - and not has_access(user.id, "write", tools.access_control, db=db) + and not AccessGrants.has_access( + user_id=user.id, + resource_type="tool", + resource_id=tools.id, + permission="write", + db=db, + ) and user.role != "admin" ): raise HTTPException( @@ -489,7 +531,13 @@ async def delete_tools_by_id( if ( tools.user_id != user.id - and not has_access(user.id, "write", tools.access_control, db=db) + and not AccessGrants.has_access( + user_id=user.id, + resource_type="tool", + resource_id=tools.id, + permission="write", + db=db, + ) and user.role != "admin" ): raise HTTPException( @@ -588,7 +636,13 @@ async def update_tools_valves_by_id( if ( tools.user_id != user.id - and not has_access(user.id, "write", tools.access_control, db=db) + and not AccessGrants.has_access( + user_id=user.id, + resource_type="tool", + resource_id=tools.id, + permission="write", + db=db, + ) and user.role != "admin" ): raise HTTPException( diff --git a/backend/open_webui/routers/users.py b/backend/open_webui/routers/users.py index 20b69bcdf7..6eca1fcac2 100644 --- a/backend/open_webui/routers/users.py +++ b/backend/open_webui/routers/users.py @@ -19,7 +19,7 @@ from open_webui.models.users import ( UserModel, UserGroupIdsModel, UserGroupIdsListResponse, - UserInfoListResponse, + UserInfoResponse, UserInfoListResponse, UserRoleUpdateForm, UserStatus, @@ -446,7 +446,7 @@ class UserActiveResponse(UserStatus): @router.get("/{user_id}", response_model=UserActiveResponse) async def get_user_by_id( - user_id: str, user=Depends(get_verified_user), db: Session = Depends(get_session) + user_id: str, user=Depends(get_admin_user), db: Session = Depends(get_session) ): # Check if user_id is a shared chat # If it is, get the user_id from the chat @@ -478,6 +478,27 @@ async def get_user_by_id( ) +@router.get("/{user_id}/info", response_model=UserInfoResponse) +async def get_user_info_by_id( + user_id: str, user=Depends(get_verified_user), db: Session = Depends(get_session) +): + user = Users.get_user_by_id(user_id, db=db) + if user: + groups = Groups.get_groups_by_member_id(user_id, db=db) + return UserInfoResponse( + **{ + **user.model_dump(), + "groups": [{"id": group.id, "name": group.name} for group in groups], + "is_active": Users.is_user_active(user_id, db=db), + } + ) + else: + raise HTTPException( + status_code=status.HTTP_400_BAD_REQUEST, + detail=ERROR_MESSAGES.USER_NOT_FOUND, + ) + + @router.get("/{user_id}/oauth/sessions") async def get_user_oauth_sessions_by_id( user_id: str, user=Depends(get_admin_user), db: Session = Depends(get_session) diff --git a/backend/open_webui/socket/main.py b/backend/open_webui/socket/main.py index 67e04e69c3..e987f9c29d 100644 --- a/backend/open_webui/socket/main.py +++ b/backend/open_webui/socket/main.py @@ -42,7 +42,7 @@ from open_webui.utils.auth import decode_token from open_webui.socket.utils import RedisDict, RedisLock, YdocManager from open_webui.tasks import create_task, stop_item_tasks from open_webui.utils.redis import get_redis_connection -from open_webui.utils.access_control import has_access, get_users_with_access +from open_webui.models.access_grants import AccessGrants from open_webui.env import ( @@ -405,7 +405,12 @@ async def join_note(sid, data): if ( user.role != "admin" and user.id != note.user_id - and not has_access(user.id, type="read", access_control=note.access_control) + and not AccessGrants.has_access( + user_id=user.id, + resource_type="note", + resource_id=note.id, + permission="read", + ) ): log.error(f"User {user.id} does not have access to note {data['note_id']}") return @@ -467,8 +472,11 @@ async def ydoc_document_join(sid, data): if ( user.get("role") != "admin" and user.get("id") != note.user_id - and not has_access( - user.get("id"), type="read", access_control=note.access_control + and not AccessGrants.has_access( + user_id=user.get("id"), + resource_type="note", + resource_id=note.id, + permission="read", ) ): log.error( @@ -537,8 +545,11 @@ async def document_save_handler(document_id, data, user): if ( user.get("role") != "admin" and user.get("id") != note.user_id - and not has_access( - user.get("id"), type="read", access_control=note.access_control + and not AccessGrants.has_access( + user_id=user.get("id"), + resource_type="note", + resource_id=note.id, + permission="read", ) ): log.error(f"User {user.get('id')} does not have access to note {note_id}") diff --git a/backend/open_webui/tasks.py b/backend/open_webui/tasks.py index d83226ffb7..cc91cde1a0 100644 --- a/backend/open_webui/tasks.py +++ b/backend/open_webui/tasks.py @@ -183,3 +183,18 @@ async def stop_item_tasks(redis: Redis, item_id: str): return result # Return the first failure return {"status": True, "message": f"All tasks for item {item_id} stopped."} + + +async def has_active_tasks(redis, chat_id: str) -> bool: + """Check if a chat has any active tasks.""" + task_ids = await list_task_ids_by_item_id(redis, chat_id) + return len(task_ids) > 0 + + +async def get_active_chat_ids(redis, chat_ids: List[str]) -> List[str]: + """Filter a list of chat_ids to only those with active tasks.""" + active = [] + for chat_id in chat_ids: + if await has_active_tasks(redis, chat_id): + active.append(chat_id) + return active diff --git a/backend/open_webui/tools/builtin.py b/backend/open_webui/tools/builtin.py index 78d64faafc..9dd3ea1bc7 100644 --- a/backend/open_webui/tools/builtin.py +++ b/backend/open_webui/tools/builtin.py @@ -36,7 +36,7 @@ from open_webui.models.chats import Chats from open_webui.models.channels import Channels, ChannelMember, Channel from open_webui.models.messages import Messages, Message from open_webui.models.groups import Groups -from open_webui.utils.sanitize import strip_markdown_code_fences +from open_webui.utils.sanitize import sanitize_code log = logging.getLogger(__name__) @@ -371,8 +371,8 @@ async def execute_code( return json.dumps({"error": "Request context not available"}) try: - # Strip markdown fences if model included them - code = strip_markdown_code_fences(code) + # Sanitize code (strips ANSI codes and markdown fences) + code = sanitize_code(code) # Import blocked modules from config (same as middleware) from open_webui.config import CODE_INTERPRETER_BLOCKED_MODULES @@ -743,10 +743,14 @@ async def view_note( user_id = __user__.get("id") user_group_ids = [group.id for group in Groups.get_groups_by_member_id(user_id)] - from open_webui.utils.access_control import has_access + from open_webui.models.access_grants import AccessGrants - if note.user_id != user_id and not has_access( - user_id, "read", note.access_control, user_group_ids + if note.user_id != user_id and not AccessGrants.has_access( + user_id=user_id, + resource_type="note", + resource_id=note.id, + permission="read", + user_group_ids=set(user_group_ids), ): return json.dumps({"error": "Access denied"}) @@ -797,7 +801,7 @@ async def write_note( form = NoteForm( title=title, data={"content": {"md": content}}, - access_control={}, # Private by default - only owner can access + access_grants=[], # Private by default - only owner can access ) new_note = Notes.insert_new_note(user_id, form) @@ -852,10 +856,14 @@ async def replace_note_content( user_id = __user__.get("id") user_group_ids = [group.id for group in Groups.get_groups_by_member_id(user_id)] - from open_webui.utils.access_control import has_access + from open_webui.models.access_grants import AccessGrants - if note.user_id != user_id and not has_access( - user_id, "write", note.access_control, user_group_ids + if note.user_id != user_id and not AccessGrants.has_access( + user_id=user_id, + resource_type="note", + resource_id=note.id, + permission="write", + user_group_ids=set(user_group_ids), ): return json.dumps({"error": "Write access denied"}) @@ -1532,7 +1540,7 @@ async def view_knowledge_file( try: from open_webui.models.files import Files from open_webui.models.knowledge import Knowledges - from open_webui.utils.access_control import has_access + from open_webui.models.access_grants import AccessGrants user_id = __user__.get("id") user_role = __user__.get("role", "user") @@ -1551,8 +1559,12 @@ async def view_knowledge_file( if ( user_role == "admin" or knowledge_base.user_id == user_id - or has_access( - user_id, "read", knowledge_base.access_control, user_group_ids + or AccessGrants.has_access( + user_id=user_id, + resource_type="knowledge", + resource_id=knowledge_base.id, + permission="read", + user_group_ids=set(user_group_ids), ) ): has_knowledge_access = True @@ -1631,7 +1643,7 @@ async def query_knowledge_files( from open_webui.models.files import Files from open_webui.models.notes import Notes from open_webui.retrieval.utils import query_collection - from open_webui.utils.access_control import has_access + from open_webui.models.access_grants import AccessGrants user_id = __user__.get("id") user_role = __user__.get("role", "user") @@ -1656,8 +1668,12 @@ async def query_knowledge_files( if knowledge and ( user_role == "admin" or knowledge.user_id == user_id - or has_access( - user_id, "read", knowledge.access_control, user_group_ids + or AccessGrants.has_access( + user_id=user_id, + resource_type="knowledge", + resource_id=knowledge.id, + permission="read", + user_group_ids=set(user_group_ids), ) ): collection_names.append(item_id) @@ -1674,7 +1690,12 @@ async def query_knowledge_files( if note and ( user_role == "admin" or note.user_id == user_id - or has_access(user_id, "read", note.access_control) + or AccessGrants.has_access( + user_id=user_id, + resource_type="note", + resource_id=note.id, + permission="read", + ) ): content = note.data.get("content", {}).get("md", "") note_results.append( @@ -1693,8 +1714,12 @@ async def query_knowledge_files( if knowledge and ( user_role == "admin" or knowledge.user_id == user_id - or has_access( - user_id, "read", knowledge.access_control, user_group_ids + or AccessGrants.has_access( + user_id=user_id, + resource_type="knowledge", + resource_id=knowledge.id, + permission="read", + user_group_ids=set(user_group_ids), ) ): collection_names.append(knowledge_id) diff --git a/backend/open_webui/utils/headers.py b/backend/open_webui/utils/headers.py index 3caee50334..f0b13c00d3 100644 --- a/backend/open_webui/utils/headers.py +++ b/backend/open_webui/utils/headers.py @@ -1,11 +1,18 @@ from urllib.parse import quote +from open_webui.env import ( + FORWARD_USER_INFO_HEADER_USER_NAME, + FORWARD_USER_INFO_HEADER_USER_ID, + FORWARD_USER_INFO_HEADER_USER_EMAIL, + FORWARD_USER_INFO_HEADER_USER_ROLE, +) + def include_user_info_headers(headers, user): return { **headers, - "X-OpenWebUI-User-Name": quote(user.name, safe=" "), - "X-OpenWebUI-User-Id": user.id, - "X-OpenWebUI-User-Email": user.email, - "X-OpenWebUI-User-Role": user.role, + FORWARD_USER_INFO_HEADER_USER_NAME: quote(user.name, safe=" "), + FORWARD_USER_INFO_HEADER_USER_ID: user.id, + FORWARD_USER_INFO_HEADER_USER_EMAIL: user.email, + FORWARD_USER_INFO_HEADER_USER_ROLE: user.role, } diff --git a/backend/open_webui/utils/middleware.py b/backend/open_webui/utils/middleware.py index 3fcaf0fe97..58ea3f249f 100644 --- a/backend/open_webui/utils/middleware.py +++ b/backend/open_webui/utils/middleware.py @@ -73,7 +73,7 @@ from open_webui.models.models import Models from open_webui.retrieval.utils import get_sources_from_items -from open_webui.utils.sanitize import strip_markdown_code_fences +from open_webui.utils.sanitize import sanitize_code from open_webui.utils.chat import generate_chat_completion from open_webui.utils.task import ( get_task_model_id, @@ -127,7 +127,10 @@ from open_webui.env import ( ENABLE_REALTIME_CHAT_SAVE, ENABLE_QUERIES_CACHE, RAG_SYSTEM_CONTEXT, + ENABLE_FORWARD_USER_INFO_HEADERS, + FORWARD_SESSION_INFO_HEADER_CHAT_ID, ) +from open_webui.utils.headers import include_user_info_headers from open_webui.constants import TASKS @@ -2207,6 +2210,12 @@ async def process_chat_payload(request, form_data, user, metadata, model): for key, value in connection_headers.items(): headers[key] = value + # Add user info headers if enabled + if ENABLE_FORWARD_USER_INFO_HEADERS and user: + headers = include_user_info_headers(headers, user) + if metadata and metadata.get("chat_id"): + headers[FORWARD_SESSION_INFO_HEADER_CHAT_ID] = metadata.get("chat_id") + mcp_clients[server_id] = MCPClient() await mcp_clients[server_id].connect( url=mcp_server_connection.get("url", ""), @@ -2768,8 +2777,21 @@ async def non_streaming_chat_response_handler(response, ctx): title = Chats.get_chat_title_by_id(metadata["chat_id"]) - # Use output from backend if provided (OR-compliant backends) + # Use output from backend if provided (OR-compliant backends), + # otherwise generate from response content response_output = response_data.get("output") + if not response_output: + response_output = [ + { + "type": "message", + "id": output_id("msg"), + "status": "completed", + "role": "assistant", + "content": [ + {"type": "output_text", "text": content} + ], + } + ] await event_emitter( { @@ -2777,11 +2799,7 @@ async def non_streaming_chat_response_handler(response, ctx): "data": { "done": True, "content": content, - **( - {"output": response_output} - if response_output - else {} - ), + "output": response_output, "title": title, }, } @@ -2794,7 +2812,7 @@ async def non_streaming_chat_response_handler(response, ctx): { "role": "assistant", "content": content, - **({"output": response_output} if response_output else {}), + "output": response_output, }, ) @@ -2869,326 +2887,54 @@ async def streaming_chat_response_handler(response, ctx): # Handle as a background task async def response_handler(response, events): - def serialize_content_blocks(content_blocks, raw=False): - content = "" - - for block in content_blocks: - if block["type"] == "text": - block_content = block["content"].strip() - if block_content: - content = f"{content}{block_content}\n" - elif block["type"] == "tool_calls": - attributes = block.get("attributes", {}) - - tool_calls = block.get("content", []) - results = block.get("results", []) - - if content and not content.endswith("\n"): - content += "\n" - - if results: - - tool_calls_display_content = "" - for tool_call in tool_calls: - - tool_call_id = tool_call.get("id", "") - tool_name = tool_call.get("function", {}).get( - "name", "" - ) - tool_arguments = tool_call.get("function", {}).get( - "arguments", "" - ) - - tool_result = None - tool_result_files = None - for result in results: - if tool_call_id == result.get("tool_call_id", ""): - tool_result = result.get("content", None) - tool_result_files = result.get("files", None) - break - - if tool_result is not None: - tool_result_embeds = result.get("embeds", "") - tool_calls_display_content = f'{tool_calls_display_content}
\nTool Executed\n
\n' - else: - tool_calls_display_content = f'{tool_calls_display_content}
\nExecuting...\n
\n' - - if not raw: - content = f"{content}{tool_calls_display_content}" - else: - tool_calls_display_content = "" - - for tool_call in tool_calls: - tool_call_id = tool_call.get("id", "") - tool_name = tool_call.get("function", {}).get( - "name", "" - ) - tool_arguments = tool_call.get("function", {}).get( - "arguments", "" - ) - - tool_calls_display_content = f'{tool_calls_display_content}\n
\nExecuting...\n
\n' - - if not raw: - content = f"{content}{tool_calls_display_content}" - - elif block["type"] == "reasoning": - reasoning_display_content = html.escape( - "\n".join( - (f"> {line}" if not line.startswith(">") else line) - for line in block["content"].splitlines() - ) - ) - - reasoning_duration = block.get("duration", None) - - start_tag = block.get("start_tag", "") - end_tag = block.get("end_tag", "") - - if content and not content.endswith("\n"): - content += "\n" - - if reasoning_duration is not None: - if raw: - content = ( - f'{content}{start_tag}{block["content"]}{end_tag}\n' - ) - else: - content = f'{content}
\nThought for {reasoning_duration} seconds\n{reasoning_display_content}\n
\n' - else: - if raw: - content = ( - f'{content}{start_tag}{block["content"]}{end_tag}\n' - ) - else: - content = f'{content}
\nThinking…\n{reasoning_display_content}\n
\n' - - elif block["type"] == "code_interpreter": - attributes = block.get("attributes", {}) - output = block.get("output", None) - lang = attributes.get("lang", "") - - content_stripped, original_whitespace = ( - split_content_and_whitespace(content) - ) - if is_opening_code_block(content_stripped): - # Remove trailing backticks that would open a new block - content = ( - content_stripped.rstrip("`").rstrip() - + original_whitespace - ) - else: - # Keep content as is - either closing backticks or no backticks - content = content_stripped + original_whitespace - - if content and not content.endswith("\n"): - content += "\n" - - if output: - output = html.escape(json.dumps(output)) - - if raw: - content = f'{content}\n{block["content"]}\n\n```output\n{output}\n```\n' - else: - content = f'{content}
\nAnalyzed\n```{lang}\n{block["content"]}\n```\n
\n' - else: - if raw: - content = f'{content}\n{block["content"]}\n\n' - else: - content = f'{content}
\nAnalyzing...\n```{lang}\n{block["content"]}\n```\n
\n' - - else: - block_content = str(block["content"]).strip() - if block_content: - content = f"{content}{block['type']}: {block_content}\n" - - return content.strip() - - return content.strip() - - def convert_content_blocks_to_messages(content_blocks, raw=False): - messages = [] - - temp_blocks = [] - for idx, block in enumerate(content_blocks): - if block["type"] == "tool_calls": - messages.append( - { - "role": "assistant", - "content": serialize_content_blocks(temp_blocks, raw), - "tool_calls": block.get("content"), - } - ) - - results = block.get("results", []) - - for result in results: - messages.append( - { - "role": "tool", - "tool_call_id": result["tool_call_id"], - "content": result.get("content", "") or "", - } - ) - temp_blocks = [] - else: - temp_blocks.append(block) - - if temp_blocks: - content = serialize_content_blocks(temp_blocks, raw) - if content: - messages.append( - { - "role": "assistant", - "content": content, - } - ) - - return messages - - def convert_content_blocks_to_output(content_blocks): + def tag_output_handler(content_type, tags, content, output): """ - Convert content_blocks to Open Responses-aligned output items. - See: https://openresponses.org/specification + Detect special tags (reasoning, solution, code_interpreter) in streaming + content and create corresponding OR-aligned output items directly. + Operates on output items instead of content_blocks. """ - output_items = [] - - def next_id(prefix): - return f"{prefix}_{uuid4().hex[:24]}" - - for block in content_blocks: - block_type = block.get("type", "") - # Use backend-provided ID if available, fallback to generated - block_id = block.get("id") - - if block_type == "text": - text_content = block.get("content", "").strip() - if text_content: - output_items.append( - { - "type": "message", - "id": block_id or next_id("msg"), - "status": "completed", - "role": "assistant", - "content": [ - {"type": "output_text", "text": text_content} - ], - } - ) - - elif block_type == "tool_calls": - tool_calls = block.get("content", []) - results = block.get("results", []) - - # Emit function_call items - for tool_call in tool_calls: - call_id = tool_call.get("id", "") - func = tool_call.get("function", {}) - output_items.append( - { - "type": "function_call", - "id": call_id - or next_id( - "fc" - ), # Use call_id as item id if available - "call_id": call_id, - "name": func.get("name", ""), - "arguments": func.get("arguments", "{}"), - "status": "completed" if results else "in_progress", - } - ) - - # Emit function_call_output items - for result in results: - output_items.append( - { - "type": "function_call_output", - "id": result.get("id") or next_id("fco"), - "call_id": result.get("tool_call_id", ""), - "output": [ - { - "type": "input_text", - "text": result.get("content", ""), - } - ], - "status": "completed", - **( - {"files": result.get("files")} - if result.get("files") - else {} - ), - **( - {"embeds": result.get("embeds")} - if result.get("embeds") - else {} - ), - } - ) - - elif block_type == "reasoning": - reasoning_content = block.get("content", "").strip() - duration = block.get("duration") - output_items.append( - { - "type": "reasoning", - "id": block_id or next_id("r"), - "status": ( - "completed" - if duration is not None - else "in_progress" - ), - "content": ( - [{"type": "output_text", "text": reasoning_content}] - if reasoning_content - else None - ), - "summary": None, - } - ) - - elif block_type == "code_interpreter": - code = block.get("content", "") - output_val = block.get("output") - attrs = block.get("attributes", {}) - output_items.append( - { - "type": "open_webui:code_interpreter", - "id": block_id or next_id("ci"), - "status": ( - "completed" - if output_val is not None - else "in_progress" - ), - "lang": attrs.get("lang", ""), - "code": code, - "output": output_val, - } - ) - - return output_items - - def tag_content_handler(content_type, tags, content, content_blocks): end_flag = False def extract_attributes(tag_content): """Extract attributes from a tag if they exist.""" attributes = {} - if not tag_content: # Ensure tag_content is not None + if not tag_content: return attributes - # Match attributes in the format: key="value" (ignores single quotes for simplicity) matches = re.findall(r'(\w+)\s*=\s*"([^"]+)"', tag_content) for key, value in matches: attributes[key] = value return attributes - if content_blocks[-1]["type"] == "text": + def get_last_text(out): + """Get text from last message item, or empty string.""" + if out and out[-1].get("type") == "message": + parts = out[-1].get("content", []) + if parts and parts[-1].get("type") == "output_text": + return parts[-1].get("text", "") + return "" + + def set_last_text(out, text): + """Set text on last message item's output_text.""" + if out and out[-1].get("type") == "message": + parts = out[-1].get("content", []) + if parts and parts[-1].get("type") == "output_text": + parts[-1]["text"] = text + + # Map content_type to output item type + output_type_map = { + "reasoning": "reasoning", + "solution": "message", # solution tags just produce text + "code_interpreter": "open_webui:code_interpreter", + } + output_item_type = output_type_map.get(content_type, content_type) + + last_type = output[-1].get("type", "") if output else "" + + if last_type == "message": for start_tag, end_tag in tags: start_tag_pattern = rf"{re.escape(start_tag)}" if start_tag.startswith("<") and start_tag.endswith(">"): - # Match start tag e.g., or - # remove both '<' and '>' from start_tag - # Match start tag with attributes start_tag_pattern = ( rf"<{re.escape(start_tag[1:-1])}(\s.*?)?>" ) @@ -3198,70 +2944,128 @@ async def streaming_chat_response_handler(response, ctx): try: attr_content = ( match.group(1) if match.group(1) else "" - ) # Ensure it's not None + ) except: attr_content = "" - attributes = extract_attributes( - attr_content - ) # Extract attributes safely + attributes = extract_attributes(attr_content) - # Capture everything before and after the matched tag - before_tag = content[ - : match.start() - ] # Content before opening tag - after_tag = content[ - match.end() : - ] # Content after opening tag + before_tag = content[: match.start()] + after_tag = content[match.end() :] - # Remove the start tag and after from the currently handling text block - content_blocks[-1]["content"] = content_blocks[-1][ - "content" - ].replace(match.group(0) + after_tag, "") - - if before_tag: - content_blocks[-1]["content"] = before_tag - - if not content_blocks[-1]["content"]: - content_blocks.pop() - - # Append the new block - content_blocks.append( - { - "type": content_type, - "start_tag": start_tag, - "end_tag": end_tag, - "attributes": attributes, - "content": "", - "started_at": time.time(), - } + # Remove the start tag and everything after from last message + current_text = get_last_text(output) + set_last_text( + output, + current_text.replace(match.group(0) + after_tag, "") ) + if before_tag: + set_last_text(output, before_tag) + + if not get_last_text(output).strip(): + # Remove empty message item + if output and output[-1].get("type") == "message": + output.pop() + + # Append the new output item + if output_item_type == "reasoning": + output.append( + { + "type": "reasoning", + "id": output_id("r"), + "status": "in_progress", + "start_tag": start_tag, + "end_tag": end_tag, + "attributes": attributes, + "content": [], + "summary": None, + "started_at": time.time(), + } + ) + elif output_item_type == "open_webui:code_interpreter": + output.append( + { + "type": "open_webui:code_interpreter", + "id": output_id("ci"), + "status": "in_progress", + "start_tag": start_tag, + "end_tag": end_tag, + "attributes": attributes, + "lang": attributes.get("lang", "python"), + "code": "", + "output": None, + "started_at": time.time(), + } + ) + else: + # solution or other text-producing tag + output.append( + { + "type": "message", + "id": output_id("msg"), + "status": "in_progress", + "role": "assistant", + "content": [{"type": "output_text", "text": ""}], + "_tag_type": content_type, + "start_tag": start_tag, + "end_tag": end_tag, + "attributes": attributes, + "started_at": time.time(), + } + ) + if after_tag: - content_blocks[-1]["content"] = after_tag - tag_content_handler( - content_type, tags, after_tag, content_blocks + # Set the after_tag content on the new item + if output_item_type == "reasoning": + output[-1]["content"] = [ + {"type": "output_text", "text": after_tag} + ] + elif output_item_type == "open_webui:code_interpreter": + output[-1]["code"] = after_tag + else: + set_last_text(output, after_tag) + + tag_output_handler( + content_type, tags, after_tag, output ) break - elif content_blocks[-1]["type"] == content_type: - start_tag = content_blocks[-1]["start_tag"] - end_tag = content_blocks[-1]["end_tag"] + + elif ( + (last_type == "reasoning" and content_type == "reasoning") + or (last_type == "open_webui:code_interpreter" and content_type == "code_interpreter") + or (last_type == "message" and output[-1].get("_tag_type") == content_type) + ): + item = output[-1] + start_tag = item.get("start_tag", "") + end_tag = item.get("end_tag", "") if end_tag.startswith("<") and end_tag.endswith(">"): - # Match end tag e.g., end_tag_pattern = rf"{re.escape(end_tag)}" else: - # Handle cases where end_tag is just a tag name end_tag_pattern = rf"{re.escape(end_tag)}" - # Check if the content has the end tag if re.search(end_tag_pattern, content): end_flag = True - block_content = content_blocks[-1]["content"] - # Strip start and end tags from the content - start_tag_pattern = rf"<{re.escape(start_tag)}(.*?)>" + # Get the block content + if last_type == "reasoning": + parts = item.get("content", []) + block_content = "" + if parts and parts[-1].get("type") == "output_text": + block_content = parts[-1].get("text", "") + elif last_type == "open_webui:code_interpreter": + block_content = item.get("code", "") + else: + block_content = get_last_text(output) + + # Strip start and end tags from content + start_tag_pattern = rf"{re.escape(start_tag)}" + if start_tag.startswith("<") and start_tag.endswith(">"): + start_tag_pattern = ( + rf"<{re.escape(start_tag[1:-1])}(\s.*?)?>" + ) block_content = re.sub( start_tag_pattern, "", block_content ).strip() @@ -3269,79 +3073,98 @@ async def streaming_chat_response_handler(response, ctx): end_tag_regex = re.compile(end_tag_pattern, re.DOTALL) split_content = end_tag_regex.split(block_content, maxsplit=1) - # Content inside the tag block_content = ( split_content[0].strip() if split_content else "" ) - - # Leftover content (everything after ``) leftover_content = ( split_content[1].strip() if len(split_content) > 1 else "" ) if block_content: - content_blocks[-1]["content"] = block_content - content_blocks[-1]["ended_at"] = time.time() - content_blocks[-1]["duration"] = int( - content_blocks[-1]["ended_at"] - - content_blocks[-1]["started_at"] - ) + # Update the item with final content + if last_type == "reasoning": + item["content"] = [ + {"type": "output_text", "text": block_content} + ] + item["ended_at"] = time.time() + item["duration"] = int( + item["ended_at"] - item["started_at"] + ) + item["status"] = "completed" + elif last_type == "open_webui:code_interpreter": + item["code"] = block_content + item["ended_at"] = time.time() + item["duration"] = int( + item["ended_at"] - item["started_at"] + ) + else: + set_last_text(output, block_content) + item["ended_at"] = time.time() - # Reset the content_blocks by appending a new text block + # Reset by appending a new message item for leftover if content_type != "code_interpreter": - if leftover_content: - - content_blocks.append( - { - "type": "text", - "content": leftover_content, - } - ) - else: - content_blocks.append( - { - "type": "text", - "content": "", - } - ) - - else: - # Remove the block if content is empty - content_blocks.pop() - - if leftover_content: - content_blocks.append( + output.append( { - "type": "text", - "content": leftover_content, + "type": "message", + "id": output_id("msg"), + "status": "in_progress", + "role": "assistant", + "content": [ + { + "type": "output_text", + "text": leftover_content, + } + ], } ) else: - content_blocks.append( + output.append( { - "type": "text", - "content": "", + "type": "message", + "id": output_id("msg"), + "status": "in_progress", + "role": "assistant", + "content": [ + { + "type": "output_text", + "text": leftover_content, + } + ], } ) + else: + # Remove the block if content is empty + output.pop() + output.append( + { + "type": "message", + "id": output_id("msg"), + "status": "in_progress", + "role": "assistant", + "content": [ + { + "type": "output_text", + "text": leftover_content, + } + ], + } + ) # Clean processed content - start_tag_pattern = rf"{re.escape(start_tag)}" + start_tag_clean = rf"{re.escape(start_tag)}" if start_tag.startswith("<") and start_tag.endswith(">"): - # Match start tag e.g., or - # remove both '<' and '>' from start_tag - # Match start tag with attributes - start_tag_pattern = ( + start_tag_clean = ( rf"<{re.escape(start_tag[1:-1])}(\s.*?)?>" ) content = re.sub( - rf"{start_tag_pattern}(.|\n)*?{re.escape(end_tag)}", + rf"{start_tag_clean}(.|\n)*?{re.escape(end_tag)}", "", content, flags=re.DOTALL, ) - return content, content_blocks, end_flag + return content, output, end_flag message = Chats.get_message_by_id_and_message_id( metadata["chat_id"], metadata["message_id"] @@ -3383,13 +3206,7 @@ async def streaming_chat_response_handler(response, ctx): else: output = [] - # Keep content_blocks for backward compatibility during transition - content_blocks = [ - { - "type": "text", - "content": content, - } - ] + usage = None reasoning_tags_param = metadata.get("params", {}).get("reasoning_tags") @@ -3430,7 +3247,6 @@ async def streaming_chat_response_handler(response, ctx): async def stream_body_handler(response, form_data): nonlocal content - nonlocal content_blocks nonlocal usage nonlocal output @@ -3668,19 +3484,26 @@ async def streaming_chat_response_handler(response, ctx): # Flush any pending text first await flush_pending_delta_data() - pending_content_blocks = content_blocks + [ - { - "type": "tool_calls", - "content": response_tool_calls, - "pending": True, - } - ] + # Build pending function_call output items for display + pending_fc_items = [] + for tc in response_tool_calls: + call_id = tc.get("id", "") + func = tc.get("function", {}) + pending_fc_items.append({ + "type": "function_call", + "id": call_id or output_id("fc"), + "call_id": call_id, + "name": func.get("name", ""), + "arguments": func.get("arguments", "{}"), + "status": "in_progress", + }) + pending_output = output + pending_fc_items await event_emitter( { "type": "chat:completion", "data": { - "content": serialize_content_blocks( - pending_content_blocks + "content": serialize_output( + pending_output ), }, } @@ -3715,52 +3538,64 @@ async def streaming_chat_response_handler(response, ctx): ) if reasoning_content: if ( - not content_blocks - or content_blocks[-1]["type"] != "reasoning" + not output + or output[-1].get("type") != "reasoning" ): - reasoning_block = { + reasoning_item = { "type": "reasoning", + "id": output_id("r"), + "status": "in_progress", "start_tag": "", "end_tag": "", "attributes": { "type": "reasoning_content" }, - "content": "", + "content": [], + "summary": None, "started_at": time.time(), } - content_blocks.append(reasoning_block) + output.append(reasoning_item) else: - reasoning_block = content_blocks[-1] + reasoning_item = output[-1] - reasoning_block["content"] += reasoning_content + # Append to reasoning content + parts = reasoning_item.get("content", []) + if parts and parts[-1].get("type") == "output_text": + parts[-1]["text"] += reasoning_content + else: + reasoning_item["content"] = [ + {"type": "output_text", "text": reasoning_content} + ] data = { - "content": serialize_content_blocks( - content_blocks - ) + "content": serialize_output(output) } if value: if ( - content_blocks - and content_blocks[-1]["type"] + output + and output[-1].get("type") == "reasoning" - and content_blocks[-1] + and output[-1] .get("attributes", {}) .get("type") == "reasoning_content" ): - reasoning_block = content_blocks[-1] - reasoning_block["ended_at"] = time.time() - reasoning_block["duration"] = int( - reasoning_block["ended_at"] - - reasoning_block["started_at"] + reasoning_item = output[-1] + reasoning_item["ended_at"] = time.time() + reasoning_item["duration"] = int( + reasoning_item["ended_at"] + - reasoning_item["started_at"] ) + reasoning_item["status"] = "completed" - content_blocks.append( + output.append( { - "type": "text", - "content": "", + "type": "message", + "id": output_id("msg"), + "status": "in_progress", + "role": "assistant", + "content": [{"type": "output_text", "text": ""}], } ) @@ -3780,44 +3615,55 @@ async def streaming_chat_response_handler(response, ctx): ) content = f"{content}{value}" - if not content_blocks: - content_blocks.append( + if ( + not output + or output[-1].get("type") != "message" + ): + output.append( { - "type": "text", - "content": "", + "type": "message", + "id": output_id("msg"), + "status": "in_progress", + "role": "assistant", + "content": [{"type": "output_text", "text": ""}], } ) - content_blocks[-1]["content"] = ( - content_blocks[-1]["content"] + value - ) + # Append value to last message item's text + msg_parts = output[-1].get("content", []) + if msg_parts and msg_parts[-1].get("type") == "output_text": + msg_parts[-1]["text"] += value + else: + output[-1]["content"] = [ + {"type": "output_text", "text": value} + ] if DETECT_REASONING_TAGS: - content, content_blocks, _ = ( - tag_content_handler( + content, output, _ = ( + tag_output_handler( "reasoning", reasoning_tags, content, - content_blocks, + output, ) ) - content, content_blocks, _ = ( - tag_content_handler( + content, output, _ = ( + tag_output_handler( "solution", DEFAULT_SOLUTION_TAGS, content, - content_blocks, + output, ) ) if DETECT_CODE_INTERPRETER: - content, content_blocks, end = ( - tag_content_handler( + content, output, end = ( + tag_output_handler( "code_interpreter", DEFAULT_CODE_INTERPRETER_TAGS, content, - content_blocks, + output, ) ) @@ -3826,9 +3672,6 @@ async def streaming_chat_response_handler(response, ctx): if ENABLE_REALTIME_CHAT_SAVE: # Save message in the database - output = convert_content_blocks_to_output( - content_blocks - ) Chats.upsert_message_to_chat_by_id_and_message_id( metadata["chat_id"], metadata["message_id"], @@ -3839,8 +3682,8 @@ async def streaming_chat_response_handler(response, ctx): ) else: data = { - "content": serialize_content_blocks( - content_blocks + "content": serialize_output( + output ), } @@ -3865,32 +3708,36 @@ async def streaming_chat_response_handler(response, ctx): continue await flush_pending_delta_data() - if content_blocks: - # Clean up the last text block - if content_blocks[-1]["type"] == "text": - content_blocks[-1]["content"] = content_blocks[-1][ - "content" - ].strip() + if output: + # Clean up the last message item + if output[-1].get("type") == "message": + parts = output[-1].get("content", []) + if parts and parts[-1].get("type") == "output_text": + parts[-1]["text"] = parts[-1]["text"].strip() - if not content_blocks[-1]["content"]: - content_blocks.pop() + if not parts[-1]["text"]: + output.pop() - if not content_blocks: - content_blocks.append( - { - "type": "text", - "content": "", - } - ) + if not output: + output.append( + { + "type": "message", + "id": output_id("msg"), + "status": "in_progress", + "role": "assistant", + "content": [{"type": "output_text", "text": ""}], + } + ) - if content_blocks[-1]["type"] == "reasoning": - reasoning_block = content_blocks[-1] - if reasoning_block.get("ended_at") is None: - reasoning_block["ended_at"] = time.time() - reasoning_block["duration"] = int( - reasoning_block["ended_at"] - - reasoning_block["started_at"] + if output[-1].get("type") == "reasoning": + reasoning_item = output[-1] + if reasoning_item.get("ended_at") is None: + reasoning_item["ended_at"] = time.time() + reasoning_item["duration"] = int( + reasoning_item["ended_at"] + - reasoning_item["started_at"] ) + reasoning_item["status"] = "completed" if response_tool_calls: tool_calls.append(response_tool_calls) @@ -3912,14 +3759,19 @@ async def streaming_chat_response_handler(response, ctx): response_tool_calls = tool_calls.pop(0) - content_blocks.append( - { - "type": "tool_calls", - "content": response_tool_calls, - } - ) + # Append function_call items for each tool call + for tc in response_tool_calls: + call_id = tc.get("id", "") + func = tc.get("function", {}) + output.append({ + "type": "function_call", + "id": call_id or output_id("fc"), + "call_id": call_id, + "name": func.get("name", ""), + "arguments": func.get("arguments", "{}"), + "status": "in_progress", + }) - output = convert_content_blocks_to_output(content_blocks) await event_emitter( { "type": "chat:completion", @@ -3955,11 +3807,7 @@ async def streaming_chat_response_handler(response, ctx): f"Error parsing tool call arguments: {tool_args}" ) - # Mutate the original tool call response params as they are passed back to the passed - # back to the LLM via the content blocks. If they are in a json block and are invalid json, - # this can cause downstream LLM integrations to fail (e.g. bedrock gateway) where response - # params are not valid json. - # Main case so far is no args = "" = invalid json. + # Ensure arguments are valid JSON for downstream LLM integrations log.debug( f"Parsed args from {tool_args} to {tool_function_params}" ) @@ -4076,11 +3924,49 @@ async def streaming_chat_response_handler(response, ctx): } ) - content_blocks[-1]["results"] = results - content_blocks.append( + # Update function_call statuses and append function_call_output items + for tc in response_tool_calls: + call_id = tc.get("id", "") + # Mark function_call as completed + for item in output: + if item.get("type") == "function_call" and item.get("call_id") == call_id: + item["status"] = "completed" + # Update arguments with parsed/sanitized version + item["arguments"] = tc.get("function", {}).get("arguments", "{}") + break + + for result in results: + output.append({ + "type": "function_call_output", + "id": output_id("fco"), + "call_id": result.get("tool_call_id", ""), + "output": [ + { + "type": "input_text", + "text": result.get("content", ""), + } + ], + "status": "completed", + **( + {"files": result.get("files")} + if result.get("files") + else {} + ), + **( + {"embeds": result.get("embeds")} + if result.get("embeds") + else {} + ), + }) + + # Append a new empty message item for the next response + output.append( { - "type": "text", - "content": "", + "type": "message", + "id": output_id("msg"), + "status": "in_progress", + "role": "assistant", + "content": [{"type": "output_text", "text": ""}], } ) @@ -4100,7 +3986,6 @@ async def streaming_chat_response_handler(response, ctx): ) tool_call_sources.clear() - output = convert_content_blocks_to_output(content_blocks) await event_emitter( { "type": "chat:completion", @@ -4118,9 +4003,7 @@ async def streaming_chat_response_handler(response, ctx): "stream": True, "messages": [ *form_data["messages"], - *convert_content_blocks_to_messages( - content_blocks, True - ), + *convert_output_to_messages(output, raw=True), ], } @@ -4144,11 +4027,11 @@ async def streaming_chat_response_handler(response, ctx): retries = 0 while ( - content_blocks[-1]["type"] == "code_interpreter" + output + and output[-1].get("type") == "open_webui:code_interpreter" and retries < MAX_RETRIES ): - output = convert_content_blocks_to_output(content_blocks) await event_emitter( { "type": "chat:completion", @@ -4162,12 +4045,13 @@ async def streaming_chat_response_handler(response, ctx): retries += 1 log.debug(f"Attempt count: {retries}") - output = "" + ci_item = output[-1] + ci_output = "" try: - if content_blocks[-1]["attributes"].get("type") == "code": - code = content_blocks[-1]["content"] - # Strip markdown fences if model included them - code = strip_markdown_code_fences(code) + if ci_item.get("attributes", {}).get("type") == "code": + code = ci_item.get("code", "") + # Sanitize code (strips ANSI codes and markdown fences) + code = sanitize_code(code) if CODE_INTERPRETER_BLOCKED_MODULES: blocking_code = textwrap.dedent( @@ -4195,7 +4079,7 @@ async def streaming_chat_response_handler(response, ctx): request.app.state.config.CODE_INTERPRETER_ENGINE == "pyodide" ): - output = await event_caller( + ci_output = await event_caller( { "type": "execute:python", "data": { @@ -4211,7 +4095,7 @@ async def streaming_chat_response_handler(response, ctx): request.app.state.config.CODE_INTERPRETER_ENGINE == "jupyter" ): - output = await execute_code_jupyter( + ci_output = await execute_code_jupyter( request.app.state.config.CODE_INTERPRETER_JUPYTER_URL, code, ( @@ -4229,14 +4113,14 @@ async def streaming_chat_response_handler(response, ctx): request.app.state.config.CODE_INTERPRETER_JUPYTER_TIMEOUT, ) else: - output = { + ci_output = { "stdout": "Code interpreter engine not configured." } - log.debug(f"Code interpreter output: {output}") + log.debug(f"Code interpreter output: {ci_output}") - if isinstance(output, dict): - stdout = output.get("stdout", "") + if isinstance(ci_output, dict): + stdout = ci_output.get("stdout", "") if isinstance(stdout, str): stdoutLines = stdout.split("\n") @@ -4254,9 +4138,9 @@ async def streaming_chat_response_handler(response, ctx): f"![Output Image]({image_url})" ) - output["stdout"] = "\n".join(stdoutLines) + ci_output["stdout"] = "\n".join(stdoutLines) - result = output.get("result", "") + result = ci_output.get("result", "") if isinstance(result, str): resultLines = result.split("\n") @@ -4271,20 +4155,23 @@ async def streaming_chat_response_handler(response, ctx): resultLines[idx] = ( f"![Output Image]({image_url})" ) - output["result"] = "\n".join(resultLines) + ci_output["result"] = "\n".join(resultLines) except Exception as e: - output = str(e) + ci_output = str(e) - content_blocks[-1]["output"] = output + ci_item["output"] = ci_output + ci_item["status"] = "completed" - content_blocks.append( + output.append( { - "type": "text", - "content": "", + "type": "message", + "id": output_id("msg"), + "status": "in_progress", + "role": "assistant", + "content": [{"type": "output_text", "text": ""}], } ) - output = convert_content_blocks_to_output(content_blocks) await event_emitter( { "type": "chat:completion", @@ -4302,12 +4189,7 @@ async def streaming_chat_response_handler(response, ctx): "stream": True, "messages": [ *form_data["messages"], - { - "role": "assistant", - "content": serialize_content_blocks( - content_blocks, raw=True - ), - }, + *convert_output_to_messages(output, raw=True), ], } @@ -4326,8 +4208,12 @@ async def streaming_chat_response_handler(response, ctx): log.debug(e) break + # Mark all in-progress items as completed + for item in output: + if item.get("status") == "in_progress": + item["status"] = "completed" + title = Chats.get_chat_title_by_id(metadata["chat_id"]) - output = convert_content_blocks_to_output(content_blocks) data = { "done": True, "content": serialize_output(output), @@ -4383,7 +4269,6 @@ async def streaming_chat_response_handler(response, ctx): if not ENABLE_REALTIME_CHAT_SAVE: # Save message in the database - output = convert_content_blocks_to_output(content_blocks) Chats.upsert_message_to_chat_by_id_and_message_id( metadata["chat_id"], metadata["message_id"], diff --git a/backend/open_webui/utils/misc.py b/backend/open_webui/utils/misc.py index b931476ca9..b2b10bf56e 100644 --- a/backend/open_webui/utils/misc.py +++ b/backend/open_webui/utils/misc.py @@ -128,22 +128,40 @@ def get_content_from_message(message: dict) -> Optional[str]: return None -def convert_output_to_messages(output: list) -> list[dict]: +def convert_output_to_messages(output: list, raw: bool = False) -> list[dict]: """ - Convert OR-aligned output items to OpenAI-format messages for LLM consumption. - - This is the inverse of convert_content_blocks_to_output() in middleware.py. + Convert OR-aligned output items to OpenAI Chat Completion-format messages. + + This reconstructs the full conversation from the stored Responses API-native + output items, including assistant messages with tool_calls arrays and tool + role messages. + + Args: + output: List of OR-aligned output items (Responses API format). + raw: If True, include reasoning blocks (with original tags) and code + interpreter blocks for LLM re-processing follow-ups. """ if not output or not isinstance(output, list): return [] - + messages = [] pending_tool_calls = [] pending_content = [] - + + def flush_pending(): + nonlocal pending_content, pending_tool_calls + if pending_content or pending_tool_calls: + messages.append({ + "role": "assistant", + "content": "\n".join(pending_content) if pending_content else "", + **({"tool_calls": pending_tool_calls} if pending_tool_calls else {}), + }) + pending_content = [] + pending_tool_calls = [] + for item in output: item_type = item.get("type", "") - + if item_type == "message": # Extract text from output_text content parts content_parts = item.get("content", []) @@ -153,58 +171,86 @@ def convert_output_to_messages(output: list) -> list[dict]: text += part.get("text", "") if text: pending_content.append(text) - + elif item_type == "function_call": # Collect tool calls to batch into assistant message + arguments = item.get("arguments", "{}") + # Ensure arguments is always a JSON string + if not isinstance(arguments, str): + arguments = json.dumps(arguments) pending_tool_calls.append({ "id": item.get("call_id", ""), "type": "function", "function": { "name": item.get("name", ""), - "arguments": item.get("arguments", "{}"), + "arguments": arguments, } }) - + elif item_type == "function_call_output": # Flush any pending content/tool_calls before adding tool result - if pending_content or pending_tool_calls: - messages.append({ - "role": "assistant", - "content": "\n".join(pending_content) if pending_content else "", - **({"tool_calls": pending_tool_calls} if pending_tool_calls else {}), - }) - pending_content = [] - pending_tool_calls = [] - + flush_pending() + # Extract text from output content parts output_parts = item.get("output", []) content = "" for part in output_parts: if part.get("type") == "input_text": content += part.get("text", "") - + messages.append({ "role": "tool", "tool_call_id": item.get("call_id", ""), "content": content, }) - + elif item_type == "reasoning": - # Skip reasoning blocks for LLM messages - pass - + if raw: + # Include reasoning with original tags for LLM re-processing + reasoning_text = "" + source_list = item.get("summary", []) or item.get("content", []) + for part in source_list: + if part.get("type") == "output_text": + reasoning_text += part.get("text", "") + elif "text" in part: + reasoning_text += part.get("text", "") + + if reasoning_text: + start_tag = item.get("start_tag", "") + end_tag = item.get("end_tag", "") + pending_content.append( + f"{start_tag}{reasoning_text}{end_tag}" + ) + # else: skip reasoning blocks for normal LLM messages + + elif item_type == "open_webui:code_interpreter": + if raw: + # Include code interpreter content for LLM re-processing + code = item.get("code", "") + code_output = item.get("output", "") + + if code: + lang = item.get("lang", "python") + pending_content.append(f"```{lang}\n{code}\n```") + + if code_output: + if isinstance(code_output, dict): + stdout = code_output.get("stdout", "") + result = code_output.get("result", "") + output_text = stdout or result + else: + output_text = str(code_output) + if output_text: + pending_content.append(f"Output:\n{output_text}") + # else: skip extension types + elif item_type.startswith("open_webui:"): - # Skip extension types + # Skip other extension types pass - + # Flush remaining content/tool_calls - if pending_content or pending_tool_calls: - messages.append({ - "role": "assistant", - "content": "\n".join(pending_content) if pending_content else "", - **({"tool_calls": pending_tool_calls} if pending_tool_calls else {}), - }) - + flush_pending() + return messages diff --git a/backend/open_webui/utils/models.py b/backend/open_webui/utils/models.py index b3a332adee..4224605f19 100644 --- a/backend/open_webui/utils/models.py +++ b/backend/open_webui/utils/models.py @@ -13,6 +13,7 @@ from open_webui.functions import get_function_models from open_webui.models.functions import Functions from open_webui.models.models import Models +from open_webui.models.access_grants import AccessGrants from open_webui.models.groups import Groups @@ -354,8 +355,12 @@ def check_model_access(user, model, db=None): raise Exception("Model not found") elif not ( user.id == model_info.user_id - or has_access( - user.id, type="read", access_control=model_info.access_control, db=db + or AccessGrants.has_access( + user_id=user.id, + resource_type="model", + resource_id=model_info.id, + permission="read", + db=db, ) ): raise Exception("Model not found") @@ -395,11 +400,13 @@ def get_filtered_models(models, user, db=None): if ( (user.role == "admin" and BYPASS_ADMIN_ACCESS_CONTROL) or user.id == model_info.user_id - or has_access( - user.id, - type="read", - access_control=model_info.access_control, + or AccessGrants.has_access( + user_id=user.id, + resource_type="model", + resource_id=model_info.id, + permission="read", user_group_ids=user_group_ids, + db=db, ) ): filtered_models.append(model) diff --git a/backend/open_webui/utils/sanitize.py b/backend/open_webui/utils/sanitize.py index 5755a9be48..344a7e08e2 100644 --- a/backend/open_webui/utils/sanitize.py +++ b/backend/open_webui/utils/sanitize.py @@ -1,5 +1,25 @@ import re +# ANSI escape code pattern - matches all common ANSI sequences +# This includes color codes, cursor movement, and other terminal control sequences +ANSI_ESCAPE_PATTERN = re.compile(r'\x1b\[[0-9;]*[A-Za-z]|\x1b\([AB]|\x1b[PX^_].*?\x1b\\|\x1b\].*?(?:\x07|\x1b\\)') + + +def strip_ansi_codes(text: str) -> str: + """ + Strip ANSI escape codes from text. + + ANSI escape codes can be introduced by LLMs that include terminal + color codes in their output. These codes cause syntax errors when + the code is sent to Jupyter for execution. + + Common ANSI codes include: + - Color codes: \x1b[31m (red), \x1b[32m (green), etc. + - Reset codes: \x1b[0m, \x1b[39m + - Cursor movement: \x1b[1A, \x1b[2J, etc. + """ + return ANSI_ESCAPE_PATTERN.sub('', text) + def strip_markdown_code_fences(code: str) -> str: """ @@ -19,3 +39,20 @@ def strip_markdown_code_fences(code: str) -> str: # Remove closing fence code = re.sub(r"\n?```\s*$", "", code) return code.strip() + + +def sanitize_code(code: str) -> str: + """ + Sanitize code for execution by applying all necessary cleanup steps. + + This is the recommended function to use before sending code to + interpreters like Jupyter or Pyodide. + + Steps applied: + 1. Strip ANSI escape codes (from LLM output) + 2. Strip markdown code fences (if model included them) + """ + code = strip_ansi_codes(code) + code = strip_markdown_code_fences(code) + return code + diff --git a/backend/open_webui/utils/tools.py b/backend/open_webui/utils/tools.py index b2f95a8297..5bb523f836 100644 --- a/backend/open_webui/utils/tools.py +++ b/backend/open_webui/utils/tools.py @@ -38,6 +38,7 @@ from open_webui.utils.misc import is_string_allowed from open_webui.models.tools import Tools from open_webui.models.users import UserModel from open_webui.models.groups import Groups +from open_webui.models.access_grants import AccessGrants from open_webui.utils.plugin import load_tool_module_by_id from open_webui.utils.access_control import has_access from open_webui.config import BYPASS_ADMIN_ACCESS_CONTROL @@ -45,7 +46,10 @@ from open_webui.env import ( AIOHTTP_CLIENT_TIMEOUT, AIOHTTP_CLIENT_TIMEOUT_TOOL_SERVER_DATA, AIOHTTP_CLIENT_SESSION_TOOL_SERVER_SSL, + ENABLE_FORWARD_USER_INFO_HEADERS, + FORWARD_SESSION_INFO_HEADER_CHAT_ID, ) +from open_webui.utils.headers import include_user_info_headers from open_webui.tools.builtin import ( search_web, fetch_url, @@ -165,7 +169,13 @@ async def get_tools( if ( not (user.role == "admin" and BYPASS_ADMIN_ACCESS_CONTROL) and tool.user_id != user.id - and not has_access(user.id, "read", tool.access_control, user_group_ids) + and not AccessGrants.has_access( + user_id=user.id, + resource_type="tool", + resource_id=tool.id, + permission="read", + user_group_ids=user_group_ids, + ) ): log.warning(f"Access denied to tool {tool_id} for user {user.id}") continue @@ -335,6 +345,13 @@ async def get_tools( for key, value in connection_headers.items(): headers[key] = value + # Add user info headers if enabled + if ENABLE_FORWARD_USER_INFO_HEADERS and user: + headers = include_user_info_headers(headers, user) + metadata = extra_params.get("__metadata__", {}) + if metadata and metadata.get("chat_id"): + headers[FORWARD_SESSION_INFO_HEADER_CHAT_ID] = metadata.get("chat_id") + def make_tool_function( function_name, tool_server_data, headers ): @@ -697,9 +714,10 @@ def convert_openapi_to_tool_payload(openapi_spec): "parameters": {"type": "object", "properties": {}, "required": []}, } - # Extract path and query parameters for param in operation.get("parameters", []): - param_name = param["name"] + param_name = param.get("name") + if not param_name: + continue param_schema = param.get("schema", {}) description = param_schema.get("description", "") if not description: @@ -971,8 +989,10 @@ async def execute_tool_server( body_params = {} for param in operation.get("parameters", []): - param_name = param["name"] - param_in = param["in"] + param_name = param.get("name") + if not param_name: + continue + param_in = param.get("in") if param_name in params: if param_in == "path": path_params[param_name] = params[param_name] diff --git a/backend/requirements-min.txt b/backend/requirements-min.txt index c4daedc446..b5cf822c5d 100644 --- a/backend/requirements-min.txt +++ b/backend/requirements-min.txt @@ -1,19 +1,19 @@ # Minimal requirements for backend to run # WIP: use this as a reference to build a minimal docker image -fastapi==0.128.0 +fastapi==0.128.5 uvicorn[standard]==0.40.0 pydantic==2.12.5 python-multipart==0.0.22 itsdangerous==2.2.0 -python-socketio==5.16.0 +python-socketio==5.16.1 python-jose==3.5.0 cryptography bcrypt==5.0.0 argon2-cffi==25.1.0 -PyJWT[crypto]==2.10.1 -authlib==1.6.6 +PyJWT[crypto]==2.11.0 +authlib==1.6.7 requests==2.32.5 aiohttp==3.13.2 # do not update to 3.13.3 - broken @@ -29,19 +29,19 @@ alembic==1.18.3 peewee==3.19.0 peewee-migrate==1.14.3 -pycrdt==0.12.45 +pycrdt==0.12.46 redis APScheduler==3.11.2 RestrictedPython==8.1 loguru==0.7.3 -asgiref==3.11.0 +asgiref==3.11.1 mcp==1.26.0 openai -langchain==1.2.7 +langchain==1.2.9 langchain-community==0.4.1 langchain-classic==1.0.1 langchain-text-splitters==1.1.0 diff --git a/backend/requirements.txt b/backend/requirements.txt index c9fdd5d0d7..957fb6ae27 100644 --- a/backend/requirements.txt +++ b/backend/requirements.txt @@ -1,16 +1,16 @@ -fastapi==0.128.0 +fastapi==0.128.5 uvicorn[standard]==0.40.0 pydantic==2.12.5 python-multipart==0.0.22 itsdangerous==2.2.0 -python-socketio==5.16.0 +python-socketio==5.16.1 python-jose==3.5.0 cryptography bcrypt==5.0.0 argon2-cffi==25.1.0 -PyJWT[crypto]==2.10.1 -authlib==1.6.6 +PyJWT[crypto]==2.11.0 +authlib==1.6.7 requests==2.32.5 aiohttp==3.13.2 # do not update to 3.13.3 - broken @@ -27,14 +27,15 @@ alembic==1.18.3 peewee==3.19.0 peewee-migrate==1.14.3 -pycrdt==0.12.45 +pycrdt==0.12.46 redis APScheduler==3.11.2 RestrictedPython==8.1 +pytz==2025.2 loguru==0.7.3 -asgiref==3.11.0 +asgiref==3.11.1 # AI libraries tiktoken @@ -42,9 +43,9 @@ mcp==1.26.0 openai anthropic -google-genai==1.60.0 +google-genai==1.62.0 -langchain==1.2.7 +langchain==1.2.9 langchain-community==0.4.1 langchain-classic==1.0.1 langchain-text-splitters==1.1.0 @@ -54,7 +55,7 @@ chromadb==1.4.1 weaviate-client==4.19.2 opensearch-py==3.1.0 -transformers==4.57.6 +transformers==5.1.0 sentence-transformers==5.2.2 accelerate pyarrow==20.0.0 # fix: pin pyarrow version to 20 for rpi compatibility #15897 @@ -62,17 +63,17 @@ einops==0.8.2 ftfy==6.3.1 chardet==5.2.0 -pypdf==6.6.2 +pypdf==6.7.0 fpdf2==2.8.5 pymdown-extensions==10.20.1 docx2txt==0.9 python-pptx==1.0.2 unstructured==0.18.31 -msoffcrypto-tool==5.4.2 +msoffcrypto-tool==6.0.0 nltk==3.9.2 Markdown==3.10.1 pypandoc==1.16.2 -pandas==2.3.3 +pandas==3.0.0 openpyxl==3.1.5 pyxlsb==1.0.10 xlrd==2.0.2 @@ -82,11 +83,11 @@ sentencepiece soundfile==0.13.1 pillow==12.1.0 -opencv-python-headless==4.13.0.90 +opencv-python-headless==4.13.0.92 rapidocr-onnxruntime==1.4.4 rank-bm25==0.2.2 -onnxruntime==1.23.2 +onnxruntime==1.24.1 faster-whisper==1.2.1 black==26.1.0 @@ -107,7 +108,7 @@ google-auth-httplib2 google-auth-oauthlib googleapis-common-protos==1.72.0 -google-cloud-storage==3.8.0 +google-cloud-storage==3.9.0 ## Databases pymongo @@ -115,12 +116,12 @@ psycopg2-binary==2.9.11 pgvector==0.4.2 PyMySQL==1.1.2 -boto3==1.42.38 +boto3==1.42.44 pymilvus==2.6.8 qdrant-client==1.16.2 playwright==1.58.0 # Caution: version must match docker-compose.playwright.yaml - Update the docker-compose.yaml if necessary -elasticsearch==9.2.1 +elasticsearch==9.3.0 pinecone==6.0.2 oracledb==3.4.2 diff --git a/pyproject.toml b/pyproject.toml index 3f07770850..eeda46681f 100644 --- a/pyproject.toml +++ b/pyproject.toml @@ -6,19 +6,19 @@ authors = [ ] license = { file = "LICENSE" } dependencies = [ - "fastapi==0.128.0", + "fastapi==0.128.5", "uvicorn[standard]==0.40.0", "pydantic==2.12.5", "python-multipart==0.0.22", "itsdangerous==2.2.0", - "python-socketio==5.16.0", + "python-socketio==5.16.1", "python-jose==3.5.0", "cryptography", "bcrypt==5.0.0", "argon2-cffi==25.1.0", - "PyJWT[crypto]==2.10.1", - "authlib==1.6.6", + "PyJWT[crypto]==2.11.0", + "authlib==1.6.7", "requests==2.32.5", "aiohttp==3.13.2", # do not update to 3.13.3 - broken @@ -35,23 +35,24 @@ dependencies = [ "peewee==3.19.0", "peewee-migrate==1.14.3", - "pycrdt==0.12.45", + "pycrdt==0.12.46", "redis", + "pytz==2025.2", "APScheduler==3.11.2", "RestrictedPython==8.1", "loguru==0.7.3", - "asgiref==3.11.0", + "asgiref==3.11.1", "tiktoken", "mcp==1.26.0", "openai", "anthropic", - "google-genai==1.60.0", + "google-genai==1.62.0", - "langchain==1.2.7", + "langchain==1.2.9", "langchain-community==0.4.1", "langchain-classic==1.0.1", "langchain-text-splitters==1.1.0", @@ -60,9 +61,9 @@ dependencies = [ "chromadb==1.4.1", "opensearch-py==3.1.0", "PyMySQL==1.1.2", - "boto3==1.42.38", + "boto3==1.42.44", - "transformers==4.57.6", + "transformers==5.1.0", "sentence-transformers==5.2.2", "accelerate", "pyarrow==20.0.0", # fix: pin pyarrow version to 20 for rpi compatibility #15897 @@ -70,17 +71,17 @@ dependencies = [ "ftfy==6.3.1", "chardet==5.2.0", - "pypdf==6.6.2", + "pypdf==6.7.0", "fpdf2==2.8.5", "pymdown-extensions==10.20.1", "docx2txt==0.9", "python-pptx==1.0.2", "unstructured==0.18.31", - "msoffcrypto-tool==5.4.2", + "msoffcrypto-tool==6.0.0", "nltk==3.9.2", "Markdown==3.10.1", "pypandoc==1.16.2", - "pandas==2.3.3", + "pandas==3.0.0", "openpyxl==3.1.5", "pyxlsb==1.0.10", "xlrd==2.0.2", @@ -91,11 +92,11 @@ dependencies = [ "azure-ai-documentintelligence==1.0.2", "pillow==12.1.0", - "opencv-python-headless==4.13.0.90", + "opencv-python-headless==4.13.0.92", "rapidocr-onnxruntime==1.4.4", "rank-bm25==0.2.2", - "onnxruntime==1.23.2", + "onnxruntime==1.24.1", "faster-whisper==1.2.1", "black==26.1.0", @@ -110,7 +111,7 @@ dependencies = [ "google-auth-oauthlib", "googleapis-common-protos==1.72.0", - "google-cloud-storage==3.8.0", + "google-cloud-storage==3.9.0", "azure-identity==1.25.1", "azure-storage-blob==12.28.0", @@ -146,7 +147,7 @@ all = [ "pytest~=8.3.2", "pytest-docker~=3.2.5", "playwright==1.58.0", # Caution: version must match docker-compose.playwright.yaml - Update the docker-compose.yaml if necessary - "elasticsearch==9.2.1", + "elasticsearch==9.3.0", "qdrant-client==1.16.2", diff --git a/src/lib/apis/channels/index.ts b/src/lib/apis/channels/index.ts index 225d8cd7cf..5715c64e89 100644 --- a/src/lib/apis/channels/index.ts +++ b/src/lib/apis/channels/index.ts @@ -3,10 +3,11 @@ import { WEBUI_API_BASE_URL } from '$lib/constants'; type ChannelForm = { type?: string; name: string; - is_private?: boolean; + is_private?: boolean | null; data?: object; meta?: object; - access_control?: object; + access_grants?: object[]; + group_ids?: string[]; user_ids?: string[]; }; diff --git a/src/lib/apis/groups/index.ts b/src/lib/apis/groups/index.ts index a74c61b83d..6089a6023f 100644 --- a/src/lib/apis/groups/index.ts +++ b/src/lib/apis/groups/index.ts @@ -99,6 +99,38 @@ export const getGroupById = async (token: string, id: string) => { return res; }; +export const getGroupInfoById = async (token: string, id: string) => { + let error = null; + + const res = await fetch(`${WEBUI_API_BASE_URL}/groups/id/${id}/info`, { + method: 'GET', + headers: { + Accept: 'application/json', + 'Content-Type': 'application/json', + authorization: `Bearer ${token}` + } + }) + .then(async (res) => { + if (!res.ok) throw await res.json(); + return res.json(); + }) + .then((json) => { + return json; + }) + .catch((err) => { + error = err.detail; + + console.error(err); + return null; + }); + + if (error) { + throw error; + } + + return res; +}; + export const updateGroupById = async (token: string, id: string, group: object) => { let error = null; diff --git a/src/lib/apis/index.ts b/src/lib/apis/index.ts index 8f35fbf881..92c28e1d9d 100644 --- a/src/lib/apis/index.ts +++ b/src/lib/apis/index.ts @@ -449,8 +449,9 @@ export const executeToolServer = async ( if (operation.parameters) { operation.parameters.forEach((param: any) => { - const paramName = param.name; - const paramIn = param.in; + const paramName = param?.name; + if (!paramName) return; + const paramIn = param?.in; if (params.hasOwnProperty(paramName)) { if (paramIn === 'path') { pathParams[paramName] = params[paramName]; @@ -1674,7 +1675,7 @@ export interface ModelMeta { profile_image_url?: string; } -export interface ModelParams {} +export interface ModelParams { } export type GlobalModelConfig = ModelConfig[]; diff --git a/src/lib/apis/knowledge/index.ts b/src/lib/apis/knowledge/index.ts index dc9dd8b88a..4c7c90484c 100644 --- a/src/lib/apis/knowledge/index.ts +++ b/src/lib/apis/knowledge/index.ts @@ -4,7 +4,7 @@ export const createNewKnowledge = async ( token: string, name: string, description: string, - accessControl: null | object + accessGrants: object[] ) => { let error = null; @@ -18,7 +18,7 @@ export const createNewKnowledge = async ( body: JSON.stringify({ name: name, description: description, - access_control: accessControl + access_grants: accessGrants }) }) .then(async (res) => { @@ -248,7 +248,7 @@ type KnowledgeUpdateForm = { name?: string; description?: string; data?: object; - access_control?: null | object; + access_grants?: object[]; }; export const updateKnowledgeById = async (token: string, id: string, form: KnowledgeUpdateForm) => { @@ -265,7 +265,7 @@ export const updateKnowledgeById = async (token: string, id: string, form: Knowl name: form?.name ? form.name : undefined, description: form?.description ? form.description : undefined, data: form?.data ? form.data : undefined, - access_control: form.access_control + access_grants: form.access_grants }) }) .then(async (res) => { diff --git a/src/lib/apis/notes/index.ts b/src/lib/apis/notes/index.ts index 55f9427e0d..341ced57ec 100644 --- a/src/lib/apis/notes/index.ts +++ b/src/lib/apis/notes/index.ts @@ -5,7 +5,7 @@ type NoteItem = { title: string; data: object; meta?: null | object; - access_control?: null | object; + access_grants?: object[]; }; export const createNewNote = async (token: string, note: NoteItem) => { diff --git a/src/lib/apis/prompts/index.ts b/src/lib/apis/prompts/index.ts index e9cd6e8481..c227c9f713 100644 --- a/src/lib/apis/prompts/index.ts +++ b/src/lib/apis/prompts/index.ts @@ -7,7 +7,7 @@ type PromptItem = { content: string; data?: object | null; meta?: object | null; - access_control?: null | object; + access_grants?: object[]; version_id?: string | null; // Active version commit_message?: string | null; // For history tracking is_production?: boolean; // Whether to set new version as production @@ -23,7 +23,7 @@ type PromptHistoryItem = { command: string; data: object; meta: object; - access_control: object | null; + access_grants: object[]; }; user_id: string; commit_message: string | null; @@ -42,7 +42,7 @@ type PromptDiff = { to_snapshot: object; content_diff: string[]; name_changed: boolean; - access_control_changed: boolean; + access_grants_changed: boolean; }; export const createNewPrompt = async (token: string, prompt: PromptItem) => { @@ -611,4 +611,3 @@ export const getPromptDiff = async ( return res; }; - diff --git a/src/lib/apis/tasks/index.ts b/src/lib/apis/tasks/index.ts new file mode 100644 index 0000000000..83299b843b --- /dev/null +++ b/src/lib/apis/tasks/index.ts @@ -0,0 +1,14 @@ +import { WEBUI_API_BASE_URL } from '$lib/constants'; + +export const checkActiveChats = async (token: string, chatIds: string[]) => { + const res = await fetch(`${WEBUI_API_BASE_URL}/tasks/active/chats`, { + method: 'POST', + headers: { + 'Content-Type': 'application/json', + Authorization: `Bearer ${token}` + }, + body: JSON.stringify({ chat_ids: chatIds }) + }); + if (!res.ok) throw await res.json(); + return res.json(); +}; diff --git a/src/lib/apis/users/index.ts b/src/lib/apis/users/index.ts index d6da54bbf9..ad669c3eb9 100644 --- a/src/lib/apis/users/index.ts +++ b/src/lib/apis/users/index.ts @@ -300,10 +300,11 @@ export const updateUserSettings = async (token: string, settings: object) => { return res; }; -export const getUserById = async (token: string, userId: string) => { + +export const getUserInfoById = async (token: string, userId: string) => { let error = null; - const res = await fetch(`${WEBUI_API_BASE_URL}/users/${userId}`, { + const res = await fetch(`${WEBUI_API_BASE_URL}/users/${userId}/info`, { method: 'GET', headers: { 'Content-Type': 'application/json', diff --git a/src/lib/components/NotificationToast.svelte b/src/lib/components/NotificationToast.svelte index 1b8d9fae8b..d232009b21 100644 --- a/src/lib/components/NotificationToast.svelte +++ b/src/lib/components/NotificationToast.svelte @@ -5,6 +5,7 @@ import { marked } from 'marked'; import { createEventDispatcher, onMount } from 'svelte'; + import XMark from '$lib/components/icons/XMark.svelte'; const dispatch = createEventDispatcher(); @@ -15,6 +16,7 @@ let startX = 0, startY = 0; let moved = false; + let closeButtonElement: HTMLButtonElement; const DRAG_THRESHOLD_PX = 6; const clickHandler = () => { @@ -22,6 +24,10 @@ dispatch('closeToast'); }; + const closeHandler = () => { + dispatch('closeToast'); + }; + function onPointerDown(e: PointerEvent) { startX = e.clientX; startY = e.clientY; @@ -43,6 +49,11 @@ // Release capture if taken (e.currentTarget as HTMLElement).releasePointerCapture?.(e.pointerId); + // Skip if clicking the close button + if (closeButtonElement && (e.target === closeButtonElement || closeButtonElement.contains(e.target as Node))) { + return; + } + // Only treat as a click if there wasn't a drag if (!moved) { clickHandler(); @@ -71,7 +82,7 @@
+ + +
favicon
diff --git a/src/lib/components/admin/Users/Groups/GroupItem.svelte b/src/lib/components/admin/Users/Groups/GroupItem.svelte index 8d4a761529..944e824d1c 100644 --- a/src/lib/components/admin/Users/Groups/GroupItem.svelte +++ b/src/lib/components/admin/Users/Groups/GroupItem.svelte @@ -9,7 +9,6 @@ import Pencil from '$lib/components/icons/Pencil.svelte'; import User from '$lib/components/icons/User.svelte'; - import UserCircleSolid from '$lib/components/icons/UserCircleSolid.svelte'; import EditGroupModal from './EditGroupModal.svelte'; export let group = { @@ -70,9 +69,6 @@ }} >
-
- -
{group.name}
diff --git a/src/lib/components/channel/ChannelInfoModal.svelte b/src/lib/components/channel/ChannelInfoModal.svelte index 44094f7801..cd1ee35243 100644 --- a/src/lib/components/channel/ChannelInfoModal.svelte +++ b/src/lib/components/channel/ChannelInfoModal.svelte @@ -15,13 +15,32 @@ import AddMembersModal from './ChannelInfoModal/AddMembersModal.svelte'; export let show = false; - export let channel = null; + export let channel: any = null; export let onUpdate = () => {}; let showAddMembersModal = false; const submitHandler = async () => {}; + const hasPublicReadGrant = (grants: any) => + Array.isArray(grants) && + grants.some( + (grant) => + grant?.principal_type === 'user' && + grant?.principal_id === '*' && + grant?.permission === 'read' + ); + + const isPublicChannel = (channel: any): boolean => { + if (channel?.type === 'group') { + if (typeof channel?.is_private === 'boolean') { + return !channel.is_private; + } + return hasPublicReadGrant(channel?.access_grants); + } + return hasPublicReadGrant(channel?.access_grants); + }; + const removeMemberHandler = async (userId) => { const res = await removeMembersById(localStorage.token, channel.id, { user_ids: [userId] @@ -62,7 +81,7 @@
{:else}
- {#if channel?.type === 'group' ? !channel?.is_private : channel?.access_control === null} + {#if isPublicChannel(channel)} {:else} diff --git a/src/lib/components/channel/MessageInput/MentionList.svelte b/src/lib/components/channel/MessageInput/MentionList.svelte index 4272a5650d..7daf692042 100644 --- a/src/lib/components/channel/MessageInput/MentionList.svelte +++ b/src/lib/components/channel/MessageInput/MentionList.svelte @@ -129,6 +129,25 @@ onDestroy(() => { window.removeEventListener('keydown', keydownListener); }); + + const hasPublicReadGrant = (grants: any) => + Array.isArray(grants) && + grants.some( + (grant) => + grant?.principal_type === 'user' && + grant?.principal_id === '*' && + grant?.permission === 'read' + ); + + const isPublicChannel = (channel: any): boolean => { + if (channel?.type === 'group') { + if (typeof channel?.is_private === 'boolean') { + return !channel.is_private; + } + return hasPublicReadGrant(channel?.access_grants); + } + return hasPublicReadGrant(channel?.access_grants); + }; {#if filteredItems.length} @@ -165,7 +184,7 @@ > {#if item.type === 'channel'}
- {#if item?.data?.access_control === null} + {#if isPublicChannel(item?.data)} {:else} diff --git a/src/lib/components/channel/Messages/Message/UserStatusLinkPreview.svelte b/src/lib/components/channel/Messages/Message/UserStatusLinkPreview.svelte index 74b2029266..749ca8df10 100644 --- a/src/lib/components/channel/Messages/Message/UserStatusLinkPreview.svelte +++ b/src/lib/components/channel/Messages/Message/UserStatusLinkPreview.svelte @@ -3,7 +3,7 @@ import { LinkPreview } from 'bits-ui'; const i18n = getContext('i18n'); - import { getUserById } from '$lib/apis/users'; + import { getUserInfoById } from '$lib/apis/users'; import UserStatus from './UserStatus.svelte'; @@ -16,7 +16,7 @@ let user = null; onMount(async () => { if (id) { - user = await getUserById(localStorage.token, id).catch((error) => { + user = await getUserInfoById(localStorage.token, id).catch((error) => { console.error('Error fetching user by ID:', error); return null; }); diff --git a/src/lib/components/channel/Navbar.svelte b/src/lib/components/channel/Navbar.svelte index 6b5d7c97c2..b02193b465 100644 --- a/src/lib/components/channel/Navbar.svelte +++ b/src/lib/components/channel/Navbar.svelte @@ -26,6 +26,25 @@ let showChannelPinnedMessagesModal = false; let showChannelInfoModal = false; + const hasPublicReadGrant = (grants: any) => + Array.isArray(grants) && + grants.some( + (grant) => + grant?.principal_type === 'user' && + grant?.principal_id === '*' && + grant?.permission === 'read' + ); + + const isPublicChannel = (channel: any): boolean => { + if (channel?.type === 'group') { + if (typeof channel?.is_private === 'boolean') { + return !channel.is_private; + } + return hasPublicReadGrant(channel?.access_grants); + } + return hasPublicReadGrant(channel?.access_grants); + }; + export let channel; export let onPin = (messageId, pinned) => {}; @@ -112,7 +131,7 @@ {/if} {:else}
- {#if channel?.type === 'group' ? !channel?.is_private : channel?.access_control === null} + {#if isPublicChannel(channel)} {:else} diff --git a/src/lib/components/chat/MessageInput.svelte b/src/lib/components/chat/MessageInput.svelte index 5914e7afce..1f16526c4b 100644 --- a/src/lib/components/chat/MessageInput.svelte +++ b/src/lib/components/chat/MessageInput.svelte @@ -151,7 +151,7 @@ return { ...file, user: undefined, - access_control: undefined + access_grants: undefined }; }), selectedToolIds, @@ -1634,7 +1634,7 @@