import os import time import datetime from pathlib import Path from fastapi import FastAPI, Request, Depends, HTTPException, status from sqlalchemy.orm import Session from fastapi.middleware.cors import CORSMiddleware from fastapi.openapi.utils import get_openapi from dotenv import load_dotenv from .utils import decode_access_token, set_current_user_ctx from .database import get_db from . import models from typing import Dict, Any, List, Optional load_dotenv() from .database import engine, Base from .routers import auth, users, roles, posts, month, security_index from .routers import hardening from .routers import files from app.routers import system_groups as system_groups_router from app.routers import systems as systems_router app = FastAPI( title="ANTT Portal API", swagger_ui_parameters={"persistAuthorization": True}, ) app.add_middleware( CORSMiddleware, allow_origins=["*"], allow_credentials=True, allow_methods=["*"], allow_headers=["*"], ) # @app.middleware("http") # async def docs_ip_whitelist(request: Request, call_next): # path = request.url.path # if path.startswith("/docs") or path.startswith("/redoc") or path.startswith("/openapi.json"): # wl = [x.strip() for x in (os.getenv("DOCS_WHITELIST_IPS") or "").split(",") if x.strip()] # client_ip = request.client.host if request.client else "" # if wl and client_ip not in wl: # raise HTTPException(status_code=403, detail="IP not allowed") # return await call_next(request) def get_real_ip(request: Request) -> str: xff = request.headers.get("x-forwarded-for") if xff: # Lấy IP đầu tiên (client thật) return xff.split(",")[0].strip() x_real_ip = request.headers.get("x-real-ip") if x_real_ip: return x_real_ip.strip() if request.client: return request.client.host return "" @app.middleware("http") async def docs_ip_whitelist(request: Request, call_next): path = request.url.path if path.startswith(("/docs", "/redoc", "/openapi.json")): wl = [ ip.strip() for ip in (os.getenv("DOCS_WHITELIST_IPS") or "").split(",") if ip.strip() ] client_ip = get_real_ip(request) if wl and client_ip not in wl: raise HTTPException( status_code=403, detail=f"IP {client_ip} not allowed" ) return await call_next(request) def _activity_log_dir() -> Path: base = os.getenv("ACTIVITY_LOG_DIR", "logs") p = Path(base).resolve() p.mkdir(parents=True, exist_ok=True) return p def _activity_log_path() -> Path: now = datetime.datetime.utcnow() month_dir = _activity_log_dir() / now.strftime("%Y%m") month_dir.mkdir(parents=True, exist_ok=True) return month_dir / f"{now.strftime('%Y-%m-%d')}.txt" @app.middleware("http") async def activity_logger(request: Request, call_next): path = request.url.path if path.startswith("/auth/login"): return await call_next(request) if path.startswith("/auth/login_sso"): return await call_next(request) method = request.method.upper() qs = request.url.query or "" ip = request.client.host if request.client else "-" start = time.time() status_code = 500 user_payload = None auth_header = request.headers.get("Authorization") or "" token = None try: scheme, token = auth_header.split() if scheme.lower() != "bearer": token = None except Exception: token = None if token: try: user_payload = decode_access_token(token) except Exception: user_payload = None try: set_current_user_ctx(user_payload) except Exception: pass try: response = await call_next(request) status_code = response.status_code return response finally: try: set_current_user_ctx(None) except Exception: pass try: user_id = "-" user_email = "-" if user_payload: try: user_id = str(user_payload.get("sub") or "-") user_email = user_payload.get("email") or "-" except Exception: pass duration_ms = int((time.time() - start) * 1000) line = f"{datetime.datetime.utcnow().isoformat()} | {ip} | {user_id} | {user_email} | {method} {path}{('?' + qs) if qs else ''} | {status_code} | {duration_ms}ms\n" with open(_activity_log_path(), "a", encoding="utf-8") as f: f.write(line) except Exception: pass @app.on_event("startup") def on_startup(): # create tables if AUTO_CREATE_DB is enabled (convenience for dev only) if os.getenv("AUTO_CREATE_DB", "0") == "1": Base.metadata.create_all(bind=engine) app.include_router(auth.router) app.include_router(roles.router, prefix="/roles", tags=["roles"]) app.include_router(users.router, prefix="/users", tags=["users"]) app.include_router(posts.router, prefix="/posts", tags=["posts"]) app.include_router(month.router, prefix="/months", tags=["months"]) app.include_router(security_index.router) app.include_router(hardening.router) # Đăng ký categories from .routers import categories app.include_router(categories.router, prefix="/categories", tags=["categories"]) app.include_router(files.router, prefix="/files", tags=["files"]) from .routers import pentest from .routers import pentest_service from .routers import jira from .routers import units from .routers import overview from .routers import scorecard from .routers import email as email_router from .routers import documents as documents_router from .routers import soc_ticket from .routers import logsource app.include_router(pentest.router, prefix="/pentest", tags=["pentest"]) app.include_router(pentest_service.router) app.include_router(soc_ticket.router) app.include_router(logsource.router) from .routers import mail_history app.include_router(mail_history.router, prefix="/mail-history", tags=["mail-history"]) from .routers import import_history as import_history_router app.include_router(import_history_router.router, prefix="/import-history", tags=["import-history"]) from .routers import manage_systems as manage_systems_router app.include_router(manage_systems_router.router) # app.include_router(system_groups_router.router) app.include_router(systems_router.router) app.include_router(jira.router, prefix="/jira", tags=["jira"]) app.include_router(units.router) app.include_router(overview.router) app.include_router(scorecard.router) app.include_router(email_router.email_roles_router) app.include_router(email_router.email_users_router) app.include_router(documents_router.router) def custom_openapi(): if app.openapi_schema: return app.openapi_schema openapi_schema = get_openapi( title="ANTT Portal API", version="1.0.0", description="API Gateway for ANTT Portal", routes=app.routes, ) openapi_schema["components"]["securitySchemes"] = { "Bearer": { "type": "http", "scheme": "bearer", "bearerFormat": "JWT", "description": "Enter JWT token (no 'Bearer' prefix needed)", } } # Áp dụng security toàn cục để UI biết cần Bearer openapi_schema["security"] = [{"Bearer": []}] # Add security requirement to all paths except /auth/login if "paths" in openapi_schema: for path, path_item in openapi_schema["paths"].items(): # Bỏ qua các endpoint auth if "/auth/" in path: continue # Thêm security cho từng operation nếu chưa có for method in ["get", "post", "put", "delete", "patch"]: if method in path_item: operation = path_item[method] if "security" not in operation: operation["security"] = [{"Bearer": []}] app.openapi_schema = openapi_schema return app.openapi_schema app.openapi = custom_openapi # Public API: system groups tree via ManageSystem (levels 0/1/2), API key + IP whitelist @app.post("/sync/system-groups", tags=["system_groups"]) def system_groups_tree( request: Request, payload: Dict[str, Any], db: Session = Depends(get_db), ): cfg_key = os.getenv("SYSTEM_GROUPS_API_KEY") if not cfg_key: raise HTTPException(status_code=500, detail="SYSTEM_GROUPS_API_KEY is not configured") use_key = str(payload.get("api_key") or "") if (use_key or "").strip() != cfg_key.strip(): raise HTTPException( status_code=status.HTTP_401_UNAUTHORIZED, detail="Invalid API key", headers={"WWW-Authenticate": "ApiKey"}, ) wl = [x.strip() for x in (os.getenv("SYSTEM_GROUPS_WHITELIST_IPS") or "").split(",") if x.strip()] client_ip = request.client.host if request.client else "" if wl and client_ip not in wl: raise HTTPException(status_code=403, detail="IP not allowed") items: List[models.ManageSystem] = db.query(models.ManageSystem).order_by(models.ManageSystem.level.asc(), models.ManageSystem.id.asc()).all() type_map = { 0: "Web quản trị", 1: "Web dịch vụ", 2: "API", 3: "Mobile app", 4: "Khác", 5: "CNTT", 6: "ATTT", } nodes: Dict[int, Dict[str, Any]] = {} for it in items: tval = int(it.type) if getattr(it, "type", None) is not None else None nodes[int(it.id)] = { "id": int(it.id), "name": it.name, "level": int(it.level), "parent_id": int(it.parent_id) if getattr(it, "parent_id", None) is not None else None, "url_ip": getattr(it, "url_ip", None), "unit_id": int(it.unit_id) if getattr(it, "unit_id", None) is not None else None, "type": tval, "type_name": type_map.get(tval) if tval is not None else None, "private": bool(it.private) if getattr(it, "private", None) is not None else None, "priority_level": int(it.priority_level) if getattr(it, "priority_level", None) is not None else None, "status": bool(it.status) if getattr(it, "status", None) is not None else None, "children": [], } roots: List[Dict[str, Any]] = [] for it in items: node = nodes[int(it.id)] pid = getattr(it, "parent_id", None) if pid is None: roots.append(node) else: parent = nodes.get(int(pid)) if parent: parent["children"].append(node) else: roots.append(node) return {"message": "success", "data": roots}