diff --git a/ft-app/app/main.py b/ft-app/app/main.py index bd67ecc..fa69e9d 100644 --- a/ft-app/app/main.py +++ b/ft-app/app/main.py @@ -1,10 +1,13 @@ from contextlib import asynccontextmanager from pathlib import Path -from fastapi import FastAPI +from fastapi import FastAPI, Request +from fastapi.responses import RedirectResponse from jinja2 import Environment, FileSystemLoader -from app.database import engine, Base +from starlette.middleware.base import BaseHTTPMiddleware +from app.database import engine, Base, SessionLocal +from app.models import User from app.seed import seed -from app.routers import contracts, analysis, data_input, admin +from app.routers import contracts, analysis, data_input, admin, auth TEMPLATES_DIR = Path(__file__).parent / "templates" @@ -14,6 +17,30 @@ def setup_jinja(app: FastAPI): app.state.templates = env +class AuthMiddleware(BaseHTTPMiddleware): + async def dispatch(self, request: Request, call_next): + # Public paths + if request.url.path in ("/auth/login", "/auth/logout") or request.url.path.startswith("/auth/"): + return await call_next(request) + + # Check session + user_id = request.cookies.get("ft_session") + if user_id: + db = SessionLocal() + try: + user = db.query(User).filter(User.id == int(user_id)).first() + if user: + request.state.user = user + return await call_next(request) + except (ValueError, TypeError): + pass + finally: + db.close() + + # Redirect to login + return RedirectResponse(f"/auth/login?next={request.url.path}", status_code=303) + + @asynccontextmanager async def lifespan(app: FastAPI): setup_jinja(app) @@ -24,6 +51,9 @@ async def lifespan(app: FastAPI): app = FastAPI(title="期货量化系统", lifespan=lifespan) +app.add_middleware(AuthMiddleware) + +app.include_router(auth.router) app.include_router(contracts.router) app.include_router(analysis.router) app.include_router(data_input.router) @@ -32,5 +62,4 @@ app.include_router(admin.router) @app.get("/") def root(): - from fastapi.responses import RedirectResponse return RedirectResponse("/contracts") diff --git a/ft-app/app/models.py b/ft-app/app/models.py index 40e25a5..06214ad 100644 --- a/ft-app/app/models.py +++ b/ft-app/app/models.py @@ -1,3 +1,5 @@ +import hashlib +import os from datetime import date from sqlalchemy import String, Integer, Float, Date, ForeignKey, Boolean, UniqueConstraint from sqlalchemy.orm import Mapped, mapped_column, relationship @@ -60,3 +62,28 @@ class PositionSnapshot(Base): position: Mapped[int] = mapped_column(Integer) delta: Mapped[int] = mapped_column(Integer, default=0) avg_cost: Mapped[float] = mapped_column(Float) + + + +class User(Base): + __tablename__ = "users" + + id: Mapped[int] = mapped_column(primary_key=True) + username: Mapped[str] = mapped_column(String(50), unique=True) + password_hash: Mapped[str] = mapped_column(String(128)) + role: Mapped[str] = mapped_column(String(20), default="user") + + @staticmethod + def hash_password(password: str) -> str: + salt = os.urandom(16) + dk = hashlib.pbkdf2_hmac("sha256", password.encode(), salt, 100000) + return salt.hex() + ":" + dk.hex() + + def check_password(self, password: str) -> bool: + try: + salt_hex, dk_hex = self.password_hash.split(":") + salt = bytes.fromhex(salt_hex) + dk = hashlib.pbkdf2_hmac("sha256", password.encode(), salt, 100000) + return dk.hex() == dk_hex + except (ValueError, AttributeError): + return False \ No newline at end of file diff --git a/ft-app/app/routers/auth.py b/ft-app/app/routers/auth.py new file mode 100644 index 0000000..5d8e3ee --- /dev/null +++ b/ft-app/app/routers/auth.py @@ -0,0 +1,52 @@ +from fastapi import APIRouter, Depends, Form, Request +from fastapi.responses import HTMLResponse, RedirectResponse +from sqlalchemy.orm import Session +from app.database import get_db +from app.models import User + +router = APIRouter(prefix="/auth", tags=["auth"]) + +SESSION_COOKIE = "ft_session" + + +def get_current_user(request: Request, db: Session = Depends(get_db)) -> User | None: + user_id = request.cookies.get(SESSION_COOKIE) + if not user_id: + return None + try: + return db.query(User).filter(User.id == int(user_id)).first() + except (ValueError, TypeError): + return None + + +@router.get("/login", response_class=HTMLResponse) +def login_page(request: Request): + template = request.app.state.templates.get_template("login.html") + return HTMLResponse(template.render(request=request)) + + +@router.post("/login") +def login( + request: Request, + username: str = Form(...), + password: str = Form(...), + db: Session = Depends(get_db), +): + user = db.query(User).filter(User.username == username).first() + if not user or not user.check_password(password): + template = request.app.state.templates.get_template("login.html") + return HTMLResponse( + template.render(request=request, error="用户名或密码错误"), + status_code=401, + ) + + resp = RedirectResponse("/contracts/", status_code=303) + resp.set_cookie(SESSION_COOKIE, str(user.id), httponly=True, max_age=86400 * 7) + return resp + + +@router.get("/logout") +def logout(): + resp = RedirectResponse("/auth/login", status_code=303) + resp.delete_cookie(SESSION_COOKIE) + return resp diff --git a/ft-app/app/seed.py b/ft-app/app/seed.py index 06a00e2..4ad4a5e 100644 --- a/ft-app/app/seed.py +++ b/ft-app/app/seed.py @@ -1,7 +1,7 @@ """Seed database from existing data files. Run once manually or on first start.""" from datetime import date from app.database import engine, Base, SessionLocal -from app.models import DailyBar, PositionSnapshot, Product, Contract +from app.models import DailyBar, PositionSnapshot, Product, Contract, User from app.engine.lock_strategy import compute_amp_5d # --- Seed OHLCV data --- @@ -202,6 +202,16 @@ def seed(): db = SessionLocal() try: + # --- Seed default user --- + if not db.query(User).first(): + admin = User( + username="admin", + password_hash=User.hash_password("admin123"), + role="admin", + ) + db.add(admin) + db.flush() + # --- Seed products and contracts --- if not db.query(Product).first(): fg = Product(code="FG", name="玻璃", exchange="CZCE") diff --git a/ft-app/app/templates/base.html b/ft-app/app/templates/base.html index cbeaea6..646b753 100644 --- a/ft-app/app/templates/base.html +++ b/ft-app/app/templates/base.html @@ -206,7 +206,11 @@ diff --git a/ft-app/app/templates/login.html b/ft-app/app/templates/login.html new file mode 100644 index 0000000..45bb471 --- /dev/null +++ b/ft-app/app/templates/login.html @@ -0,0 +1,70 @@ + + + + + +登录 · 期货量化 + + + +
+

📊 期货量化系统

+

请输入账号密码

+ + {% if error %} +
{{ error }}
+ {% endif %} + +
+
+ + +
+
+ + +
+ +
+
+ +