Files
saas_backend/app/services/auth/access_service.py
T

113 lines
4.3 KiB
Python

from sqlalchemy.orm import Session, joinedload
from app.models.auth.access_model import Access
from typing import List, Any, Optional
import uuid
from app.core.redis import sync_redis_client
import json
import logging
from app.models.auth.module_access_model import ModuleAccess
from app.models.auth.module_model import Module
from app.models.auth.tenant_module_model import TenantModule
from datetime import datetime
logger = logging.getLogger(__name__)
class AccessService:
@staticmethod
def get_all_accesses(db: Session, category: str = None, tenant_id: Optional[uuid.UUID] = None) -> List[any]:
cache_key = f"saas:access:v4:active:{category if category else 'full'}:{str(tenant_id) if tenant_id else 'global'}"
cached_data = None
if cached_data:
try:
data_list = json.loads(cached_data)
class SimpleAccess:
def __init__(self, **kwargs):
for k, v in kwargs.items():
setattr(self, k, v)
deserialized_list = []
for item in data_list:
if "created_at" in item and item["created_at"]:
try:
item["created_at"] = datetime.fromisoformat(item["created_at"])
except ValueError:
item["created_at"] = None
deserialized_list.append(SimpleAccess(**item))
return deserialized_list
except Exception as e:
logger.warning(f"Access cache read error: {e}")
query = db.query(Access)
if category:
query = query.filter(Access.category == category)
if tenant_id:
query = query.filter(Access.category != "Superadmin")
saas_accesses = query.all()
for access in saas_accesses:
access.module_name = "SaaS (Internal)"
module_query = (
db.query(ModuleAccess)
.join(Module, ModuleAccess.module_id == Module.id)
.filter(Module.status == "active")
.options(joinedload(ModuleAccess.module))
)
if category:
module_query = module_query.filter(ModuleAccess.category == category)
if tenant_id:
active_tms = (
db.query(TenantModule.module_id)
.filter(
TenantModule.tenant_id == tenant_id,
TenantModule.is_active == True
)
.all()
)
active_module_ids = [tm[0] for tm in active_tms]
module_query = module_query.filter(ModuleAccess.module_id.in_(active_module_ids))
module_accesses = module_query.all()
for ma in module_accesses:
if ma.module:
ma.module_name = ma.module.module_name
result = saas_accesses + module_accesses
try:
if sync_redis_client.client:
serialized = []
for item in result:
serialized.append({
"id": str(item.id),
"access_code": item.access_code,
"name": item.name,
"category": item.category,
"parent_id": str(item.parent_id) if item.parent_id else None,
"module_id": str(item.module_id) if getattr(item, "module_id", None) else None,
"module_name": getattr(item, "module_name", None),
"created_at": item.created_at.isoformat() if item.created_at else None,
})
sync_redis_client.client.set(cache_key, json.dumps(serialized), ex=3600)
except Exception as e:
logger.warning(f"Access cache write error: {e}")
return result
@staticmethod
def get_access_categories(db: Session) -> List[str]:
categories = db.query(Access.category).distinct().all()
module_categories = db.query(ModuleAccess.category).distinct().all()
all_cats = set([cat[0] for cat in categories] + [cat[0] for cat in module_categories])
return list(all_cats)