295 lines
10 KiB
Python
295 lines
10 KiB
Python
from sqlalchemy import create_engine
|
|
from sqlalchemy.orm import sessionmaker, declarative_base, Session
|
|
from app.core.config.Config import settings
|
|
|
|
Base = declarative_base()
|
|
|
|
# 1. Fallback connection URLs
|
|
core_url = settings.CORE_DATABASE_URL or settings.DATABASE_URL
|
|
crm_url = settings.CRM_DATABASE_URL or settings.DATABASE_URL
|
|
commerce_url = settings.COMMERCE_DATABASE_URL or settings.DATABASE_URL
|
|
|
|
# 2. Engines configuration with connection pooling
|
|
engine_core = create_engine(
|
|
core_url,
|
|
future=True,
|
|
pool_pre_ping=True,
|
|
connect_args={"init_command": "SET time_zone='+00:00'"}
|
|
)
|
|
|
|
engine_crm = create_engine(
|
|
crm_url,
|
|
future=True,
|
|
pool_pre_ping=True,
|
|
connect_args={"init_command": "SET time_zone='+00:00'"}
|
|
)
|
|
|
|
engine_commerce = create_engine(
|
|
commerce_url,
|
|
future=True,
|
|
pool_pre_ping=True,
|
|
connect_args={"init_command": "SET time_zone='+00:00'"}
|
|
)
|
|
|
|
# Maintain default engine alias for backwards compatibility
|
|
engine = engine_core
|
|
|
|
# 3. Dynamic Routing Session
|
|
class RoutingSession(Session):
|
|
def get_bind(self, mapper=None, clause=None):
|
|
table_name = None
|
|
if mapper:
|
|
table_name = getattr(mapper.persist_selectable, "name", None)
|
|
elif clause is not None:
|
|
if hasattr(clause, "table"):
|
|
table_name = getattr(clause.table, "name", None)
|
|
elif hasattr(clause, "froms") and clause.froms:
|
|
table_name = getattr(clause.froms[0], "name", None)
|
|
|
|
if table_name:
|
|
# CRM Workshop database routing
|
|
if table_name in [
|
|
"service_types", "repair_services",
|
|
"repair_variants", "repair_variant_images", "parts", "part_device_compatibility",
|
|
"repair_variant_parts", "stock_movements", "purchase_orders", "purchase_order_items",
|
|
"contacts", "contact_addresses"
|
|
]:
|
|
return engine_crm
|
|
|
|
# Core Identity & Platform database routing
|
|
if table_name in [
|
|
"departments", "designations", "roles", "permissions",
|
|
"users", "user_sessions", "audit_logs",
|
|
"countries", "states", "cities", "settings", "file_uploads",
|
|
"role_permissions"
|
|
]:
|
|
return engine_core
|
|
|
|
# Storefront Commerce database routing
|
|
if table_name in [
|
|
"device_series", "device_models",
|
|
"products", "product_variants", "product_images", "variant_attributes",
|
|
"variant_images", "attribute_types", "categories", "brands", "tags",
|
|
"collections", "storefront_contents", "product_reviews", "product_review_images",
|
|
"seo_metadata", "migration_jobs", "migration_batches", "migration_errors",
|
|
"migration_snapshots", "migration_job_checkpoints", "migration_media_items",
|
|
"media_groups", "media_library", "mapping_configs"
|
|
]:
|
|
return engine_commerce
|
|
|
|
# Default engine fallback
|
|
return engine_commerce
|
|
|
|
SessionLocal = sessionmaker(
|
|
class_=RoutingSession,
|
|
autocommit=False,
|
|
autoflush=False
|
|
)
|
|
|
|
from fastapi import Request
|
|
|
|
def get_db(request: Request):
|
|
db = SessionLocal()
|
|
|
|
# Resolve Request ID
|
|
req_id = getattr(request.state, "request_id", None)
|
|
if not req_id:
|
|
import ulid
|
|
req_id = str(ulid.ULID())
|
|
request.state.request_id = req_id
|
|
|
|
db.info["request_id"] = req_id
|
|
db.info["user_id"] = None
|
|
|
|
# Resolve IP Address
|
|
x_forwarded_for = request.headers.get("x-forwarded-for")
|
|
db.info["ip_address"] = x_forwarded_for.split(",")[0].strip() if x_forwarded_for else (request.client.host if request.client else "127.0.0.1")
|
|
|
|
# Resolve User Agent
|
|
db.info["user_agent"] = request.headers.get("user-agent", "")
|
|
|
|
# Save session reference in request state so auth logic can inject user_id back to it
|
|
request.state.db_session = db
|
|
|
|
try:
|
|
yield db
|
|
finally:
|
|
db.close()
|
|
|
|
# --- Automated Audit Log Event Listener ---
|
|
from sqlalchemy import event
|
|
|
|
@event.listens_for(SessionLocal, "before_flush")
|
|
def receive_before_flush(session, flush_context, instances):
|
|
req_id = session.info.get("request_id")
|
|
if not req_id:
|
|
return
|
|
|
|
from app.models.AuditLogModel import AuditLog
|
|
import json
|
|
import ulid
|
|
|
|
def serialize_val(val):
|
|
if val is None:
|
|
return None
|
|
if hasattr(val, "isoformat"):
|
|
return val.isoformat()
|
|
if hasattr(val, "to_eng_string"):
|
|
return str(val)
|
|
if isinstance(val, (dict, list)):
|
|
return val
|
|
try:
|
|
json.dumps(val)
|
|
return val
|
|
except Exception:
|
|
return str(val)
|
|
|
|
def get_model_dict(obj):
|
|
mapper = obj.__class__.__mapper__
|
|
data = {}
|
|
for col in mapper.column_attrs:
|
|
val = getattr(obj, col.key)
|
|
data[col.key] = serialize_val(val)
|
|
return data
|
|
|
|
logs_to_add = []
|
|
|
|
# 1. New objects (Insert)
|
|
for obj in session.new:
|
|
if isinstance(obj, AuditLog):
|
|
continue
|
|
entity_type = obj.__class__.__name__
|
|
mapper = obj.__class__.__mapper__
|
|
pk_keys = [col.key for col in mapper.primary_key]
|
|
entity_id = "-".join([str(getattr(obj, k)) for k in pk_keys]) if pk_keys else "transient"
|
|
if not entity_id or entity_id == "None" or entity_id == "transient":
|
|
# Primary Key might not be flushed yet. Resolve via common ID attributes:
|
|
if hasattr(obj, "user_id") and obj.user_id:
|
|
entity_id = str(obj.user_id)
|
|
elif hasattr(obj, "product_id") and obj.product_id:
|
|
entity_id = str(obj.product_id)
|
|
elif hasattr(obj, "category_id") and obj.category_id:
|
|
entity_id = str(obj.category_id)
|
|
elif hasattr(obj, "brand_id") and obj.brand_id:
|
|
entity_id = str(obj.brand_id)
|
|
elif hasattr(obj, "order_id") and obj.order_id:
|
|
entity_id = str(obj.order_id)
|
|
else:
|
|
entity_id = "transient"
|
|
|
|
action = "create"
|
|
if entity_type == "UserSession":
|
|
action = "login"
|
|
|
|
new_val = get_model_dict(obj)
|
|
|
|
log_entry = AuditLog(
|
|
audit_id=str(ulid.ULID()),
|
|
request_id=req_id,
|
|
user_id=session.info.get("user_id"),
|
|
entity_type=entity_type,
|
|
entity_id=entity_id,
|
|
action=action,
|
|
old_value=None,
|
|
new_value=new_val,
|
|
ip_address=session.info.get("ip_address") or "127.0.0.1",
|
|
user_agent=session.info.get("user_agent")
|
|
)
|
|
logs_to_add.append(log_entry)
|
|
|
|
# 2. Dirty objects (Update)
|
|
for obj in session.dirty:
|
|
if isinstance(obj, AuditLog):
|
|
continue
|
|
if not session.is_modified(obj):
|
|
continue
|
|
|
|
entity_type = obj.__class__.__name__
|
|
mapper = obj.__class__.__mapper__
|
|
pk_keys = [col.key for col in mapper.primary_key]
|
|
entity_id = "-".join([str(getattr(obj, k)) for k in pk_keys]) if pk_keys else "transient"
|
|
if not entity_id or entity_id == "None" or entity_id == "transient":
|
|
if hasattr(obj, "user_id") and obj.user_id:
|
|
entity_id = str(obj.user_id)
|
|
elif hasattr(obj, "product_id") and obj.product_id:
|
|
entity_id = str(obj.product_id)
|
|
elif hasattr(obj, "category_id") and obj.category_id:
|
|
entity_id = str(obj.category_id)
|
|
elif hasattr(obj, "brand_id") and obj.brand_id:
|
|
entity_id = str(obj.brand_id)
|
|
elif hasattr(obj, "order_id") and obj.order_id:
|
|
entity_id = str(obj.order_id)
|
|
|
|
old_val_dict = {}
|
|
new_val_dict = {}
|
|
|
|
from sqlalchemy.orm import attributes
|
|
for col in mapper.column_attrs:
|
|
hist = attributes.get_history(obj, col.key)
|
|
if hist.has_changes():
|
|
old_v = hist.deleted[0] if hist.deleted else None
|
|
new_v = hist.added[0] if hist.added else None
|
|
old_val_dict[col.key] = serialize_val(old_v)
|
|
new_val_dict[col.key] = serialize_val(new_v)
|
|
|
|
if not old_val_dict and not new_val_dict:
|
|
continue
|
|
|
|
action = "update"
|
|
if entity_type == "Order" and "status" in new_val_dict:
|
|
action = f"order_status_{new_val_dict['status'].lower()}"
|
|
elif entity_type == "UserSession" and "is_active" in new_val_dict and not new_val_dict["is_active"]:
|
|
action = "logout"
|
|
|
|
log_entry = AuditLog(
|
|
audit_id=str(ulid.ULID()),
|
|
request_id=req_id,
|
|
user_id=session.info.get("user_id"),
|
|
entity_type=entity_type,
|
|
entity_id=entity_id,
|
|
action=action,
|
|
old_value=old_val_dict,
|
|
new_value=new_val_dict,
|
|
ip_address=session.info.get("ip_address") or "127.0.0.1",
|
|
user_agent=session.info.get("user_agent")
|
|
)
|
|
logs_to_add.append(log_entry)
|
|
|
|
# 3. Deleted objects (Delete)
|
|
for obj in session.deleted:
|
|
if isinstance(obj, AuditLog):
|
|
continue
|
|
entity_type = obj.__class__.__name__
|
|
mapper = obj.__class__.__mapper__
|
|
pk_keys = [col.key for col in mapper.primary_key]
|
|
entity_id = "-".join([str(getattr(obj, k)) for k in pk_keys]) if pk_keys else "transient"
|
|
if not entity_id or entity_id == "None" or entity_id == "transient":
|
|
if hasattr(obj, "user_id") and obj.user_id:
|
|
entity_id = str(obj.user_id)
|
|
elif hasattr(obj, "product_id") and obj.product_id:
|
|
entity_id = str(obj.product_id)
|
|
elif hasattr(obj, "category_id") and obj.category_id:
|
|
entity_id = str(obj.category_id)
|
|
elif hasattr(obj, "brand_id") and obj.brand_id:
|
|
entity_id = str(obj.brand_id)
|
|
elif hasattr(obj, "order_id") and obj.order_id:
|
|
entity_id = str(obj.order_id)
|
|
|
|
old_val = get_model_dict(obj)
|
|
|
|
log_entry = AuditLog(
|
|
audit_id=str(ulid.ULID()),
|
|
request_id=req_id,
|
|
user_id=session.info.get("user_id"),
|
|
entity_type=entity_type,
|
|
entity_id=entity_id,
|
|
action="delete",
|
|
old_value=old_val,
|
|
new_value=None,
|
|
ip_address=session.info.get("ip_address") or "127.0.0.1",
|
|
user_agent=session.info.get("user_agent")
|
|
)
|
|
logs_to_add.append(log_entry)
|
|
|
|
for log in logs_to_add:
|
|
session.add(log)
|