diff --git a/app/app.py b/app/app.py index 44036a3..bee348e 100644 --- a/app/app.py +++ b/app/app.py @@ -18,6 +18,24 @@ engine = create_engine(os.getenv("DATABASE_URL")) SessionLocal = sessionmaker(bind=engine) Base.metadata.create_all(engine) + + +def require_capability(cap): + def decorator(f): + @wraps(f) + def wrapper(*args, **kwargs): + db = SessionLocal() + email = session["user"]["email"] + u = db.query(User).filter_by(email=email).first() + db.close() + + if not u or not getattr(u, cap, False): + return jsonify({"error": "Forbidden", "missing_capability": cap}), 403 + return f(*args, **kwargs) + return wrapper + return decorator + + oauth = OAuth(app) google = oauth.register( name="google", @@ -45,9 +63,24 @@ def login(): def oauth_callback(): token = google.authorize_access_token() user_info = google.get("userinfo").json() - session["user"] = user_info + #session["user"] = user_info + db = SessionLocal() + u = db.query(User).filter_by(email=user_info["email"]).first() + if not u: + u = User( + email=user_info["email"], + name=user_info.get("name", ""), + role="user", + can_use_audio=False, + can_use_vision=False, + can_share=True, + ) + db.add(u) + db.commit() + db.close() return redirect("/") + @app.route("/api/new_conversation", methods=["POST"]) @login_required def new_conversation(): @@ -244,6 +277,7 @@ def stream(): @app.route("/api/share/", methods=["POST"]) @login_required +@require_capability("can_share") def share(cid): target_email = request.json.get("email") db = SessionLocal()