diff --git a/app/app.py b/app/app.py index b6ef5ce..4e4d1ad 100644 --- a/app/app.py +++ b/app/app.py @@ -20,6 +20,9 @@ engine = create_engine(os.getenv("DATABASE_URL")) SessionLocal = sessionmaker(bind=engine) Base.metadata.create_all(engine) +# TODO: +# conversation.user_email == email +# msgs = db.query(Message).filter_by(conversation_id=cid).all() sock = Sock(app) @@ -404,6 +407,13 @@ def login_required(f): # If using DB authentication if "db_user" not in session: return redirect("/auth/login") + + email = session["db_user"]["email"] + u = db.query(User).filter_by(email=email).first() + if not u: + session.pop("db_user", None) + return redirect("/auth/login") + return f(*args, **kwargs) return wrapper @@ -730,7 +740,7 @@ def list_conversations(): @rate_limited def export_md(cid): db = SessionLocal() - conv = db.query(Conversation).filter_by(id=cid).first() + conv = db.query(Conversation).filter_by(id=cid, user_email=email).first() msgs = conv.messages md = f"# {conv.title}\n\n" @@ -793,7 +803,7 @@ def import_file(cid): @rate_limited def export_pdf(cid): db = SessionLocal() - conv = db.query(Conversation).filter_by(id=cid).first() + conv = db.query(Conversation).filter_by(id=cid, user_email=email).first() msgs = conv.messages html = "

{}

".format(conv.title) @@ -815,7 +825,11 @@ def export_pdf(cid): @rate_limited def history(cid): db = SessionLocal() - msgs = db.query(Message).filter_by(conversation_id=cid).order_by(Message.created_at).all() + #msgs = db.query(Message).filter_by(conversation_id=cid, user_email=email).order_by(Message.created_at).all() + msgs = db.query(Message).join(Conversation).filter( + Conversation.id == cid, + Conversation.user_email == email + ).all() data = [{"role": m.role, "content": markdown(m.content)} for m in msgs] db.close() return jsonify(data) @@ -846,7 +860,7 @@ def upload_attachment(cid): @rate_limited def list_attachments(cid): db = SessionLocal() - atts = db.query(Attachment).filter_by(conversation_id=cid).all() + atts = db.query(Attachment).filter_by(conversation_id=cid, user_email=email).all() data = [{"id": a.id, "filename": a.filename} for a in atts] db.close() return jsonify(data) @@ -857,7 +871,7 @@ def list_attachments(cid): @rate_limited def download(att_id): db = SessionLocal() - att = db.query(Attachment).filter_by(id=att_id).first() + att = db.query(Attachment).filter_by(id=att_id, user_email=email).first() db.close() if not att: @@ -866,6 +880,19 @@ def download(att_id): return send_file(att.path, as_attachment=True, download_name=att.filename) + +@app.route("/api/file/") +@login_required +def get_file(file_id): + db = SessionLocal() + f = db.query(File).filter_by(id=file_id, user_email=email).first() + if not f: + return jsonify({"erro": "Arquivo não encontrado"}), 404 + return send_file(f.path) + + + + @app.route("/api/stream") @login_required @rate_limited @@ -908,7 +935,7 @@ def stream(): def share(cid): target_email = request.json.get("email") db = SessionLocal() - conv = db.query(Conversation).filter_by(id=cid).first() + conv = db.query(Conversation).filter_by(id=cid, user_email=email).first() conv.shared = target_email db.commit() db.close() @@ -921,7 +948,7 @@ def share(cid): def update_tags(cid): tags = request.json.get("tags", []) db = SessionLocal() - conv = db.query(Conversation).filter_by(id=cid).first() + conv = db.query(Conversation).filter_by(id=cid, user_email=email).first() conv.tags = ",".join(tags) db.commit() db.close()