diff --git a/config.py b/config.py index 273e130..fdb6f4d 100644 --- a/config.py +++ b/config.py @@ -11,7 +11,10 @@ class Settings(BaseSettings): boss_url: str boss_aes_key: str boss_hmac_key: str - + twitter_client_id: str + twitter_client_secret: str + twitter_redirect_uri: str + model_config = SettingsConfigDict(env_file=".env", extra="ignore") cookie_secure: bool = True webhook_url: str diff --git a/database.py b/database.py index e51e262..b78302c 100644 --- a/database.py +++ b/database.py @@ -59,6 +59,14 @@ class PlayerRank(Base): createdAt = Column(DateTime, server_default=func.now()) updatedAt = Column(DateTime, onupdate=func.now()) +class TwitterLink(Base): + __tablename__ = "twitter_link" + pid = Column(Integer, primary_key=True) + twitter_handle = Column(String) + twitter_token_enc = Column(String) + twitter_refresh_token_enc = Column(String) + miidata_enc = Column(String) + engine = create_engine( settings.db_url, pool_pre_ping=True, diff --git a/main.py b/main.py index 7874ce9..9bc3e7c 100644 --- a/main.py +++ b/main.py @@ -12,6 +12,7 @@ from routes.equipment import equipment_history from routes.equipment import equipment from services.boss_retrieval import process_boss_file from routes import Ranking +from routes import twitter_link from contextlib import asynccontextmanager from datetime import datetime, timedelta from sqlalchemy import delete @@ -155,6 +156,7 @@ app.include_router(me.router, prefix="/api/v1") app.include_router(equipment_history.router, prefix="/api/v1") app.include_router(equipment.router, prefix="/api/v1") app.include_router(Ranking.router, prefix="/api/v1") +app.include_router(twitter_link.router, prefix="/api/v1") if __name__ == '__main__': uvicorn.run("main:app", host="0.0.0.0", port=settings.port, reload=True) \ No newline at end of file diff --git a/requirements.txt b/requirements.txt index b69f878..ca0e8ed 100644 --- a/requirements.txt +++ b/requirements.txt @@ -30,4 +30,5 @@ starlette==0.52.1 typing-inspection==0.4.2 typing_extensions==4.15.0 urllib3==2.6.3 -uvicorn==0.41.0 \ No newline at end of file +uvicorn==0.41.0 +tweepy==4.16.0 \ No newline at end of file diff --git a/routes/twitter_link.py b/routes/twitter_link.py new file mode 100644 index 0000000..0de250d --- /dev/null +++ b/routes/twitter_link.py @@ -0,0 +1,133 @@ +from fastapi import APIRouter, Request, Depends, HTTPException +from sqlalchemy.orm import Session as DBSession +from database import SessionLocal, User, Session as UserSession, TwitterLink +from config import settings, cipher +from services import auth +import secrets +import hashlib +import base64 +import httpx +import json + +router = APIRouter() + +def get_db(): + db = SessionLocal() + try: + yield db + finally: + db.close() + +pkce_store = {} + +@router.get("/me/twitter/link") +async def link_twitter(request: Request, db: DBSession = Depends(get_db)): + session_id = request.cookies.get("session_id") + db_session = db.query(UserSession).filter(UserSession.id == session_id).first() + if not db_session: + raise HTTPException(status_code=401) + + state = secrets.token_urlsafe(16) + verifier = secrets.token_urlsafe(64) + + sha256_hash = hashlib.sha256(verifier.encode('utf-8')).digest() + challenge = base64.urlsafe_b64encode(sha256_hash).decode('utf-8').replace('=', '') + + pkce_store[state] = { + "verifier": verifier, + "username": db_session.username + } + + scopes = "tweet.read tweet.write users.read offline.access" + encoded_scopes = scopes.replace(" ", "%20") + + base_url = "https://twitter.com/i/oauth2/authorize" + params = [ + "response_type=code", + f"client_id={settings.twitter_client_id}", + f"redirect_uri={settings.twitter_redirect_uri}", + f"scope={encoded_scopes}", + f"state={state}", + f"code_challenge={challenge}", + "code_challenge_method=S256" + ] + + return {"url": f"{base_url}?{'&'.join(params)}"} + +@router.get("/me/twitter/confirm") +async def confirm_twitter(state: str, code: str, db: DBSession = Depends(get_db)): + stored_data = pkce_store.get(state) + if not stored_data: + raise HTTPException(status_code=400, detail="State not found.") + + auth_str = f"{settings.twitter_client_id}:{settings.twitter_client_secret}" + encoded_auth = base64.b64encode(auth_str.encode()).decode() + + async with httpx.AsyncClient() as client: + token_res = await client.post( + "https://api.twitter.com/2/oauth2/token", + headers={"Authorization": f"Basic {encoded_auth}", "Content-Type": "application/x-www-form-urlencoded"}, + data={ + "grant_type": "authorization_code", + "code": code, + "redirect_uri": settings.twitter_redirect_uri, + "code_verifier": stored_data["verifier"], + } + ) + + token_data = token_res.json() + if token_res.status_code != 200: + raise HTTPException(status_code=401, detail="Twitter auth failed") + + async with httpx.AsyncClient() as client: + user_res = await client.get( + "https://api.twitter.com/2/users/me", + headers={"Authorization": f"Bearer {token_data['access_token']}"} + ) + handle = user_res.json().get("data", {}).get("username") + + user = db.query(User).filter(User.username == stored_data["username"]).first() + if not user or not user.spfn_pass_enc: + raise HTTPException(status_code=404, detail="Local user not found") + + decrypted_pass = cipher.decrypt(user.spfn_pass_enc.encode()).decode() + profile_auth = auth.get_token(user.username, decrypted_pass) + + profile_response = auth.get_profile(profile_auth["token"]) + + if isinstance(profile_response, str): + try: + profile_data = json.loads(profile_response) + except json.JSONDecodeError: + raise HTTPException(status_code=500, detail="Failed to parse player profile JSON") + else: + profile_data = profile_response + + if not profile_data: + raise HTTPException(status_code=500, detail="Could not fetch player profile") + + pid = profile_data.get("pid") + mii_data = profile_data.get("mii", {}).get("data") + + def encrypt_val(val: str): + return cipher.encrypt(val.encode()).decode() if val else None + + tw_link = db.query(TwitterLink).filter(TwitterLink.pid == pid).first() + if not tw_link: + tw_link = TwitterLink(pid=pid) + db.add(tw_link) + + tw_link.twitter_handle = handle + tw_link.twitter_token_enc = encrypt_val(token_data["access_token"]) + tw_link.twitter_refresh_token_enc = encrypt_val(token_data.get("refresh_token")) + tw_link.miidata_enc = encrypt_val(mii_data) + + try: + db.commit() + except Exception as e: + db.rollback() + raise HTTPException(status_code=500, detail=f"Database error: {str(e)}") + + pkce_store.pop(state, None) + + return {"status": "success", "pid": pid, "handle": handle} \ No newline at end of file