refac
This commit is contained in:
@@ -45,26 +45,22 @@ PAGE_ITEM_COUNT = 30
|
||||
############################
|
||||
|
||||
|
||||
@router.get("/", response_model=list[PromptModel])
|
||||
async def get_prompts(
|
||||
user=Depends(get_verified_user), db: Session = Depends(get_session)
|
||||
):
|
||||
if user.role == "admin" and BYPASS_ADMIN_ACCESS_CONTROL:
|
||||
@router.get('/', response_model=list[PromptModel])
|
||||
async def get_prompts(user=Depends(get_verified_user), db: Session = Depends(get_session)):
|
||||
if user.role == 'admin' and BYPASS_ADMIN_ACCESS_CONTROL:
|
||||
prompts = Prompts.get_prompts(db=db)
|
||||
else:
|
||||
prompts = Prompts.get_prompts_by_user_id(user.id, "read", db=db)
|
||||
prompts = Prompts.get_prompts_by_user_id(user.id, 'read', db=db)
|
||||
|
||||
return prompts
|
||||
|
||||
|
||||
@router.get("/tags", response_model=list[str])
|
||||
async def get_prompt_tags(
|
||||
user=Depends(get_verified_user), db: Session = Depends(get_session)
|
||||
):
|
||||
if user.role == "admin" and BYPASS_ADMIN_ACCESS_CONTROL:
|
||||
@router.get('/tags', response_model=list[str])
|
||||
async def get_prompt_tags(user=Depends(get_verified_user), db: Session = Depends(get_session)):
|
||||
if user.role == 'admin' and BYPASS_ADMIN_ACCESS_CONTROL:
|
||||
return Prompts.get_tags(db=db)
|
||||
else:
|
||||
prompts = Prompts.get_prompts_by_user_id(user.id, "read", db=db)
|
||||
prompts = Prompts.get_prompts_by_user_id(user.id, 'read', db=db)
|
||||
tags = set()
|
||||
for prompt in prompts:
|
||||
if prompt.tags:
|
||||
@@ -72,7 +68,7 @@ async def get_prompt_tags(
|
||||
return sorted(list(tags))
|
||||
|
||||
|
||||
@router.get("/list", response_model=PromptAccessListResponse)
|
||||
@router.get('/list', response_model=PromptAccessListResponse)
|
||||
async def get_prompt_list(
|
||||
query: Optional[str] = None,
|
||||
view_option: Optional[str] = None,
|
||||
@@ -90,37 +86,35 @@ async def get_prompt_list(
|
||||
|
||||
filter = {}
|
||||
if query:
|
||||
filter["query"] = query
|
||||
filter['query'] = query
|
||||
if view_option:
|
||||
filter["view_option"] = view_option
|
||||
filter['view_option'] = view_option
|
||||
if tag:
|
||||
filter["tag"] = tag
|
||||
filter['tag'] = tag
|
||||
if order_by:
|
||||
filter["order_by"] = order_by
|
||||
filter['order_by'] = order_by
|
||||
if direction:
|
||||
filter["direction"] = direction
|
||||
filter['direction'] = direction
|
||||
|
||||
# Pre-fetch user group IDs once - used for both filter and write_access check
|
||||
groups = Groups.get_groups_by_member_id(user.id, db=db)
|
||||
user_group_ids = {group.id for group in groups}
|
||||
|
||||
if not (user.role == "admin" and BYPASS_ADMIN_ACCESS_CONTROL):
|
||||
if not (user.role == 'admin' and BYPASS_ADMIN_ACCESS_CONTROL):
|
||||
if groups:
|
||||
filter["group_ids"] = [group.id for group in groups]
|
||||
filter['group_ids'] = [group.id for group in groups]
|
||||
|
||||
filter["user_id"] = user.id
|
||||
filter['user_id'] = user.id
|
||||
|
||||
result = Prompts.search_prompts(
|
||||
user.id, filter=filter, skip=skip, limit=limit, db=db
|
||||
)
|
||||
result = Prompts.search_prompts(user.id, filter=filter, skip=skip, limit=limit, db=db)
|
||||
|
||||
# Batch-fetch writable prompt IDs in a single query instead of N has_access calls
|
||||
prompt_ids = [prompt.id for prompt in result.items]
|
||||
writable_prompt_ids = AccessGrants.get_accessible_resource_ids(
|
||||
user_id=user.id,
|
||||
resource_type="prompt",
|
||||
resource_type='prompt',
|
||||
resource_ids=prompt_ids,
|
||||
permission="write",
|
||||
permission='write',
|
||||
user_group_ids=user_group_ids,
|
||||
db=db,
|
||||
)
|
||||
@@ -130,7 +124,7 @@ async def get_prompt_list(
|
||||
PromptAccessResponse(
|
||||
**prompt.model_dump(),
|
||||
write_access=(
|
||||
(user.role == "admin" and BYPASS_ADMIN_ACCESS_CONTROL)
|
||||
(user.role == 'admin' and BYPASS_ADMIN_ACCESS_CONTROL)
|
||||
or user.id == prompt.user_id
|
||||
or prompt.id in writable_prompt_ids
|
||||
),
|
||||
@@ -146,23 +140,23 @@ async def get_prompt_list(
|
||||
############################
|
||||
|
||||
|
||||
@router.post("/create", response_model=Optional[PromptModel])
|
||||
@router.post('/create', response_model=Optional[PromptModel])
|
||||
async def create_new_prompt(
|
||||
request: Request,
|
||||
form_data: PromptForm,
|
||||
user=Depends(get_verified_user),
|
||||
db: Session = Depends(get_session),
|
||||
):
|
||||
if user.role != "admin" and not (
|
||||
if user.role != 'admin' and not (
|
||||
has_permission(
|
||||
user.id,
|
||||
"workspace.prompts",
|
||||
'workspace.prompts',
|
||||
request.app.state.config.USER_PERMISSIONS,
|
||||
db=db,
|
||||
)
|
||||
or has_permission(
|
||||
user.id,
|
||||
"workspace.prompts_import",
|
||||
'workspace.prompts_import',
|
||||
request.app.state.config.USER_PERMISSIONS,
|
||||
db=db,
|
||||
)
|
||||
@@ -193,34 +187,32 @@ async def create_new_prompt(
|
||||
############################
|
||||
|
||||
|
||||
@router.get("/command/{command}", response_model=Optional[PromptAccessResponse])
|
||||
async def get_prompt_by_command(
|
||||
command: str, user=Depends(get_verified_user), db: Session = Depends(get_session)
|
||||
):
|
||||
@router.get('/command/{command}', response_model=Optional[PromptAccessResponse])
|
||||
async def get_prompt_by_command(command: str, user=Depends(get_verified_user), db: Session = Depends(get_session)):
|
||||
prompt = Prompts.get_prompt_by_command(command, db=db)
|
||||
|
||||
if prompt:
|
||||
if (
|
||||
user.role == "admin"
|
||||
user.role == 'admin'
|
||||
or prompt.user_id == user.id
|
||||
or AccessGrants.has_access(
|
||||
user_id=user.id,
|
||||
resource_type="prompt",
|
||||
resource_type='prompt',
|
||||
resource_id=prompt.id,
|
||||
permission="read",
|
||||
permission='read',
|
||||
db=db,
|
||||
)
|
||||
):
|
||||
return PromptAccessResponse(
|
||||
**prompt.model_dump(),
|
||||
write_access=(
|
||||
(user.role == "admin" and BYPASS_ADMIN_ACCESS_CONTROL)
|
||||
(user.role == 'admin' and BYPASS_ADMIN_ACCESS_CONTROL)
|
||||
or user.id == prompt.user_id
|
||||
or AccessGrants.has_access(
|
||||
user_id=user.id,
|
||||
resource_type="prompt",
|
||||
resource_type='prompt',
|
||||
resource_id=prompt.id,
|
||||
permission="write",
|
||||
permission='write',
|
||||
db=db,
|
||||
)
|
||||
),
|
||||
@@ -237,34 +229,32 @@ async def get_prompt_by_command(
|
||||
############################
|
||||
|
||||
|
||||
@router.get("/id/{prompt_id}", response_model=Optional[PromptAccessResponse])
|
||||
async def get_prompt_by_id(
|
||||
prompt_id: str, user=Depends(get_verified_user), db: Session = Depends(get_session)
|
||||
):
|
||||
@router.get('/id/{prompt_id}', response_model=Optional[PromptAccessResponse])
|
||||
async def get_prompt_by_id(prompt_id: str, user=Depends(get_verified_user), db: Session = Depends(get_session)):
|
||||
prompt = Prompts.get_prompt_by_id(prompt_id, db=db)
|
||||
|
||||
if prompt:
|
||||
if (
|
||||
user.role == "admin"
|
||||
user.role == 'admin'
|
||||
or prompt.user_id == user.id
|
||||
or AccessGrants.has_access(
|
||||
user_id=user.id,
|
||||
resource_type="prompt",
|
||||
resource_type='prompt',
|
||||
resource_id=prompt.id,
|
||||
permission="read",
|
||||
permission='read',
|
||||
db=db,
|
||||
)
|
||||
):
|
||||
return PromptAccessResponse(
|
||||
**prompt.model_dump(),
|
||||
write_access=(
|
||||
(user.role == "admin" and BYPASS_ADMIN_ACCESS_CONTROL)
|
||||
(user.role == 'admin' and BYPASS_ADMIN_ACCESS_CONTROL)
|
||||
or user.id == prompt.user_id
|
||||
or AccessGrants.has_access(
|
||||
user_id=user.id,
|
||||
resource_type="prompt",
|
||||
resource_type='prompt',
|
||||
resource_id=prompt.id,
|
||||
permission="write",
|
||||
permission='write',
|
||||
db=db,
|
||||
)
|
||||
),
|
||||
@@ -281,7 +271,7 @@ async def get_prompt_by_id(
|
||||
############################
|
||||
|
||||
|
||||
@router.post("/id/{prompt_id}/update", response_model=Optional[PromptModel])
|
||||
@router.post('/id/{prompt_id}/update', response_model=Optional[PromptModel])
|
||||
async def update_prompt_by_id(
|
||||
prompt_id: str,
|
||||
form_data: PromptForm,
|
||||
@@ -301,12 +291,12 @@ async def update_prompt_by_id(
|
||||
prompt.user_id != user.id
|
||||
and not AccessGrants.has_access(
|
||||
user_id=user.id,
|
||||
resource_type="prompt",
|
||||
resource_type='prompt',
|
||||
resource_id=prompt.id,
|
||||
permission="write",
|
||||
permission='write',
|
||||
db=db,
|
||||
)
|
||||
and user.role != "admin"
|
||||
and user.role != 'admin'
|
||||
):
|
||||
raise HTTPException(
|
||||
status_code=status.HTTP_401_UNAUTHORIZED,
|
||||
@@ -338,7 +328,7 @@ async def update_prompt_by_id(
|
||||
############################
|
||||
|
||||
|
||||
@router.post("/id/{prompt_id}/update/meta", response_model=Optional[PromptModel])
|
||||
@router.post('/id/{prompt_id}/update/meta', response_model=Optional[PromptModel])
|
||||
async def update_prompt_metadata(
|
||||
prompt_id: str,
|
||||
form_data: PromptMetadataForm,
|
||||
@@ -358,12 +348,12 @@ async def update_prompt_metadata(
|
||||
prompt.user_id != user.id
|
||||
and not AccessGrants.has_access(
|
||||
user_id=user.id,
|
||||
resource_type="prompt",
|
||||
resource_type='prompt',
|
||||
resource_id=prompt.id,
|
||||
permission="write",
|
||||
permission='write',
|
||||
db=db,
|
||||
)
|
||||
and user.role != "admin"
|
||||
and user.role != 'admin'
|
||||
):
|
||||
raise HTTPException(
|
||||
status_code=status.HTTP_401_UNAUTHORIZED,
|
||||
@@ -379,9 +369,7 @@ async def update_prompt_metadata(
|
||||
detail=f"Command '/{form_data.command}' is already in use",
|
||||
)
|
||||
|
||||
updated_prompt = Prompts.update_prompt_metadata(
|
||||
prompt.id, form_data.name, form_data.command, form_data.tags, db=db
|
||||
)
|
||||
updated_prompt = Prompts.update_prompt_metadata(prompt.id, form_data.name, form_data.command, form_data.tags, db=db)
|
||||
if updated_prompt:
|
||||
return updated_prompt
|
||||
else:
|
||||
@@ -391,7 +379,7 @@ async def update_prompt_metadata(
|
||||
)
|
||||
|
||||
|
||||
@router.post("/id/{prompt_id}/update/version", response_model=Optional[PromptModel])
|
||||
@router.post('/id/{prompt_id}/update/version', response_model=Optional[PromptModel])
|
||||
async def set_prompt_version(
|
||||
prompt_id: str,
|
||||
form_data: PromptVersionUpdateForm,
|
||||
@@ -409,21 +397,19 @@ async def set_prompt_version(
|
||||
prompt.user_id != user.id
|
||||
and not AccessGrants.has_access(
|
||||
user_id=user.id,
|
||||
resource_type="prompt",
|
||||
resource_type='prompt',
|
||||
resource_id=prompt.id,
|
||||
permission="write",
|
||||
permission='write',
|
||||
db=db,
|
||||
)
|
||||
and user.role != "admin"
|
||||
and user.role != 'admin'
|
||||
):
|
||||
raise HTTPException(
|
||||
status_code=status.HTTP_401_UNAUTHORIZED,
|
||||
detail=ERROR_MESSAGES.ACCESS_PROHIBITED,
|
||||
)
|
||||
|
||||
updated_prompt = Prompts.update_prompt_version(
|
||||
prompt.id, form_data.version_id, db=db
|
||||
)
|
||||
updated_prompt = Prompts.update_prompt_version(prompt.id, form_data.version_id, db=db)
|
||||
if updated_prompt:
|
||||
return updated_prompt
|
||||
else:
|
||||
@@ -442,7 +428,7 @@ class PromptAccessGrantsForm(BaseModel):
|
||||
access_grants: list[dict]
|
||||
|
||||
|
||||
@router.post("/id/{prompt_id}/access/update", response_model=Optional[PromptModel])
|
||||
@router.post('/id/{prompt_id}/access/update', response_model=Optional[PromptModel])
|
||||
async def update_prompt_access_by_id(
|
||||
request: Request,
|
||||
prompt_id: str,
|
||||
@@ -461,12 +447,12 @@ async def update_prompt_access_by_id(
|
||||
prompt.user_id != user.id
|
||||
and not AccessGrants.has_access(
|
||||
user_id=user.id,
|
||||
resource_type="prompt",
|
||||
resource_type='prompt',
|
||||
resource_id=prompt.id,
|
||||
permission="write",
|
||||
permission='write',
|
||||
db=db,
|
||||
)
|
||||
and user.role != "admin"
|
||||
and user.role != 'admin'
|
||||
):
|
||||
raise HTTPException(
|
||||
status_code=status.HTTP_401_UNAUTHORIZED,
|
||||
@@ -478,10 +464,10 @@ async def update_prompt_access_by_id(
|
||||
user.id,
|
||||
user.role,
|
||||
form_data.access_grants,
|
||||
"sharing.public_prompts",
|
||||
'sharing.public_prompts',
|
||||
)
|
||||
|
||||
AccessGrants.set_access_grants("prompt", prompt_id, form_data.access_grants, db=db)
|
||||
AccessGrants.set_access_grants('prompt', prompt_id, form_data.access_grants, db=db)
|
||||
|
||||
return Prompts.get_prompt_by_id(prompt_id, db=db)
|
||||
|
||||
@@ -491,10 +477,8 @@ async def update_prompt_access_by_id(
|
||||
############################
|
||||
|
||||
|
||||
@router.post("/id/{prompt_id}/toggle", response_model=Optional[PromptModel])
|
||||
async def toggle_prompt_active(
|
||||
prompt_id: str, user=Depends(get_verified_user), db: Session = Depends(get_session)
|
||||
):
|
||||
@router.post('/id/{prompt_id}/toggle', response_model=Optional[PromptModel])
|
||||
async def toggle_prompt_active(prompt_id: str, user=Depends(get_verified_user), db: Session = Depends(get_session)):
|
||||
prompt = Prompts.get_prompt_by_id(prompt_id, db=db)
|
||||
|
||||
if not prompt:
|
||||
@@ -507,12 +491,12 @@ async def toggle_prompt_active(
|
||||
prompt.user_id != user.id
|
||||
and not AccessGrants.has_access(
|
||||
user_id=user.id,
|
||||
resource_type="prompt",
|
||||
resource_type='prompt',
|
||||
resource_id=prompt.id,
|
||||
permission="write",
|
||||
permission='write',
|
||||
db=db,
|
||||
)
|
||||
and user.role != "admin"
|
||||
and user.role != 'admin'
|
||||
):
|
||||
raise HTTPException(
|
||||
status_code=status.HTTP_401_UNAUTHORIZED,
|
||||
@@ -533,10 +517,8 @@ async def toggle_prompt_active(
|
||||
############################
|
||||
|
||||
|
||||
@router.delete("/id/{prompt_id}/delete", response_model=bool)
|
||||
async def delete_prompt_by_id(
|
||||
prompt_id: str, user=Depends(get_verified_user), db: Session = Depends(get_session)
|
||||
):
|
||||
@router.delete('/id/{prompt_id}/delete', response_model=bool)
|
||||
async def delete_prompt_by_id(prompt_id: str, user=Depends(get_verified_user), db: Session = Depends(get_session)):
|
||||
prompt = Prompts.get_prompt_by_id(prompt_id, db=db)
|
||||
|
||||
if not prompt:
|
||||
@@ -549,12 +531,12 @@ async def delete_prompt_by_id(
|
||||
prompt.user_id != user.id
|
||||
and not AccessGrants.has_access(
|
||||
user_id=user.id,
|
||||
resource_type="prompt",
|
||||
resource_type='prompt',
|
||||
resource_id=prompt.id,
|
||||
permission="write",
|
||||
permission='write',
|
||||
db=db,
|
||||
)
|
||||
and user.role != "admin"
|
||||
and user.role != 'admin'
|
||||
):
|
||||
raise HTTPException(
|
||||
status_code=status.HTTP_401_UNAUTHORIZED,
|
||||
@@ -570,7 +552,7 @@ async def delete_prompt_by_id(
|
||||
############################
|
||||
|
||||
|
||||
@router.get("/id/{prompt_id}/history", response_model=list[PromptHistoryResponse])
|
||||
@router.get('/id/{prompt_id}/history', response_model=list[PromptHistoryResponse])
|
||||
async def get_prompt_history(
|
||||
prompt_id: str,
|
||||
page: int = 0,
|
||||
@@ -590,13 +572,13 @@ async def get_prompt_history(
|
||||
|
||||
# Check read access
|
||||
if not (
|
||||
user.role == "admin"
|
||||
user.role == 'admin'
|
||||
or prompt.user_id == user.id
|
||||
or AccessGrants.has_access(
|
||||
user_id=user.id,
|
||||
resource_type="prompt",
|
||||
resource_type='prompt',
|
||||
resource_id=prompt.id,
|
||||
permission="read",
|
||||
permission='read',
|
||||
db=db,
|
||||
)
|
||||
):
|
||||
@@ -605,13 +587,11 @@ async def get_prompt_history(
|
||||
detail=ERROR_MESSAGES.ACCESS_PROHIBITED,
|
||||
)
|
||||
|
||||
history = PromptHistories.get_history_by_prompt_id(
|
||||
prompt.id, limit=PAGE_SIZE, offset=page * PAGE_SIZE, db=db
|
||||
)
|
||||
history = PromptHistories.get_history_by_prompt_id(prompt.id, limit=PAGE_SIZE, offset=page * PAGE_SIZE, db=db)
|
||||
return history
|
||||
|
||||
|
||||
@router.get("/id/{prompt_id}/history/{history_id}", response_model=PromptHistoryModel)
|
||||
@router.get('/id/{prompt_id}/history/{history_id}', response_model=PromptHistoryModel)
|
||||
async def get_prompt_history_entry(
|
||||
prompt_id: str,
|
||||
history_id: str,
|
||||
@@ -629,13 +609,13 @@ async def get_prompt_history_entry(
|
||||
|
||||
# Check read access
|
||||
if not (
|
||||
user.role == "admin"
|
||||
user.role == 'admin'
|
||||
or prompt.user_id == user.id
|
||||
or AccessGrants.has_access(
|
||||
user_id=user.id,
|
||||
resource_type="prompt",
|
||||
resource_type='prompt',
|
||||
resource_id=prompt.id,
|
||||
permission="read",
|
||||
permission='read',
|
||||
db=db,
|
||||
)
|
||||
):
|
||||
@@ -654,7 +634,7 @@ async def get_prompt_history_entry(
|
||||
return history_entry
|
||||
|
||||
|
||||
@router.delete("/id/{prompt_id}/history/{history_id}", response_model=bool)
|
||||
@router.delete('/id/{prompt_id}/history/{history_id}', response_model=bool)
|
||||
async def delete_prompt_history_entry(
|
||||
prompt_id: str,
|
||||
history_id: str,
|
||||
@@ -672,13 +652,13 @@ async def delete_prompt_history_entry(
|
||||
|
||||
# Check write access
|
||||
if not (
|
||||
user.role == "admin"
|
||||
user.role == 'admin'
|
||||
or prompt.user_id == user.id
|
||||
or AccessGrants.has_access(
|
||||
user_id=user.id,
|
||||
resource_type="prompt",
|
||||
resource_type='prompt',
|
||||
resource_id=prompt.id,
|
||||
permission="write",
|
||||
permission='write',
|
||||
db=db,
|
||||
)
|
||||
):
|
||||
@@ -691,7 +671,7 @@ async def delete_prompt_history_entry(
|
||||
if prompt.version_id == history_id:
|
||||
raise HTTPException(
|
||||
status_code=status.HTTP_400_BAD_REQUEST,
|
||||
detail="Cannot delete the active production version",
|
||||
detail='Cannot delete the active production version',
|
||||
)
|
||||
|
||||
success = PromptHistories.delete_history_entry(history_id, db=db)
|
||||
@@ -704,7 +684,7 @@ async def delete_prompt_history_entry(
|
||||
return success
|
||||
|
||||
|
||||
@router.get("/id/{prompt_id}/history/diff")
|
||||
@router.get('/id/{prompt_id}/history/diff')
|
||||
async def get_prompt_diff(
|
||||
prompt_id: str,
|
||||
from_id: str,
|
||||
@@ -723,13 +703,13 @@ async def get_prompt_diff(
|
||||
|
||||
# Check read access
|
||||
if not (
|
||||
user.role == "admin"
|
||||
user.role == 'admin'
|
||||
or prompt.user_id == user.id
|
||||
or AccessGrants.has_access(
|
||||
user_id=user.id,
|
||||
resource_type="prompt",
|
||||
resource_type='prompt',
|
||||
resource_id=prompt.id,
|
||||
permission="read",
|
||||
permission='read',
|
||||
db=db,
|
||||
)
|
||||
):
|
||||
@@ -742,7 +722,7 @@ async def get_prompt_diff(
|
||||
if not diff:
|
||||
raise HTTPException(
|
||||
status_code=status.HTTP_404_NOT_FOUND,
|
||||
detail="One or both history entries not found",
|
||||
detail='One or both history entries not found',
|
||||
)
|
||||
|
||||
return diff
|
||||
|
||||
Reference in New Issue
Block a user