This commit is contained in:
Timothy Jaeryang Baek
2026-02-19 16:53:21 -06:00
parent d664922feb
commit 8c485b260f
8 changed files with 271 additions and 38 deletions
+65 -10
View File
@@ -144,11 +144,20 @@ class ModelsTable:
) -> 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:
def _to_model_model(
self,
model: Model,
access_grants: Optional[list[AccessGrantModel]] = None,
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)
model_data["access_grants"] = (
access_grants
if access_grants is not None
else self._get_access_grants(model_data["id"], db=db)
)
return ModelModel.model_validate(model_data)
def insert_new_model(
@@ -181,8 +190,14 @@ class ModelsTable:
def get_all_models(self, db: Optional[Session] = None) -> list[ModelModel]:
with get_db_context(db) as db:
all_models = db.query(Model).all()
model_ids = [model.id for model in all_models]
grants_map = AccessGrants.get_grants_by_resources("model", model_ids, db=db)
return [
self._to_model_model(model, db=db) for model in db.query(Model).all()
self._to_model_model(
model, access_grants=grants_map.get(model.id, []), db=db
)
for model in all_models
]
def get_models(self, db: Optional[Session] = None) -> list[ModelUserResponse]:
@@ -190,9 +205,11 @@ class ModelsTable:
all_models = db.query(Model).filter(Model.base_model_id != None).all()
user_ids = list(set(model.user_id for model in all_models))
model_ids = [model.id for model in all_models]
users = Users.get_users_by_user_ids(user_ids, db=db) if user_ids else []
users_dict = {user.id: user for user in users}
grants_map = AccessGrants.get_grants_by_resources("model", model_ids, db=db)
models = []
for model in all_models:
@@ -200,7 +217,11 @@ class ModelsTable:
models.append(
ModelUserResponse.model_validate(
{
**self._to_model_model(model, db=db).model_dump(),
**self._to_model_model(
model,
access_grants=grants_map.get(model.id, []),
db=db,
).model_dump(),
"user": user.model_dump() if user else None,
}
)
@@ -209,9 +230,16 @@ class ModelsTable:
def get_base_models(self, db: Optional[Session] = None) -> list[ModelModel]:
with get_db_context(db) as db:
all_models = (
db.query(Model).filter(Model.base_model_id == None).all()
)
model_ids = [model.id for model in all_models]
grants_map = AccessGrants.get_grants_by_resources("model", model_ids, db=db)
return [
self._to_model_model(model, db=db)
for model in db.query(Model).filter(Model.base_model_id == None).all()
self._to_model_model(
model, access_grants=grants_map.get(model.id, []), db=db
)
for model in all_models
]
def get_models_by_user_id(
@@ -325,11 +353,18 @@ class ModelsTable:
items = query.all()
model_ids = [model.id for model, _ in items]
grants_map = AccessGrants.get_grants_by_resources("model", model_ids, db=db)
models = []
for model, user in items:
models.append(
ModelUserResponse(
**self._to_model_model(model, db=db).model_dump(),
**self._to_model_model(
model,
access_grants=grants_map.get(model.id, []),
db=db,
).model_dump(),
user=(
UserResponse(**UserModel.model_validate(user).model_dump())
if user
@@ -356,7 +391,18 @@ class ModelsTable:
try:
with get_db_context(db) as db:
models = db.query(Model).filter(Model.id.in_(ids)).all()
return [self._to_model_model(model, db=db) for model in models]
model_ids = [model.id for model in models]
grants_map = AccessGrants.get_grants_by_resources(
"model", model_ids, db=db
)
return [
self._to_model_model(
model,
access_grants=grants_map.get(model.id, []),
db=db,
)
for model in models
]
except Exception:
return []
@@ -465,9 +511,18 @@ class ModelsTable:
db.commit()
all_models = db.query(Model).all()
model_ids = [model.id for model in all_models]
grants_map = AccessGrants.get_grants_by_resources(
"model", model_ids, db=db
)
return [
self._to_model_model(model, db=db)
for model in db.query(Model).all()
self._to_model_model(
model,
access_grants=grants_map.get(model.id, []),
db=db,
)
for model in all_models
]
except Exception as e:
log.exception(f"Error syncing models for user {user_id}: {e}")