This commit is contained in:
Timothy Jaeryang Baek
2026-02-08 21:24:20 -06:00
parent 42763cbbd8
commit 0f78451c2b
8 changed files with 398 additions and 444 deletions
+51 -58
View File
@@ -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
-5
View File
@@ -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):
+84 -41
View File
@@ -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()
+70 -89
View File
@@ -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}")
+42 -138
View File
@@ -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
+31 -19
View File
@@ -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
+76 -54
View File
@@ -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()
+44 -40
View File
@@ -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()