static / app.py
WACATW's picture
app.py
98cccaa verified
Raw History Blame Contribute Delete
5.1 kB
import os
import json
import gradio as gr
from huggingface_hub import HfApi, upload_file, download_file, list_repo_files, delete_file, snapshot_download
# 读取环境变量
HF_TOKEN = os.getenv("HF_TOKEN")
DATASET_REPO = os.getenv("DATASET_REPO")
api = HfApi(token=HF_TOKEN)
# 仓库内存储路径
USER_DB_PATH = "user_db.json"
# 加载用户账号数据
def load_user_db():
try:
download_file(repo_id=DATASET_REPO, filename=USER_DB_PATH, local_dir="./tmp", repo_type="dataset")
with open(f"./tmp/{USER_DB_PATH}", "r", encoding="utf-8") as f:
return json.load(f)
except:
return {}
# 保存用户账号
def save_user_db(db):
with open(f"./tmp/{USER_DB_PATH}", "w", encoding="utf-8") as f:
json.dump(db, f, ensure_ascii=False)
upload_file(path_or_fileobj=f"./tmp/{USER_DB_PATH}", path_in_repo=USER_DB_PATH, repo_id=DATASET_REPO, repo_type="dataset)
# 注册
def register(username, pwd):
db = load_user_db()
if username in db:
return "用户名已存在"
db[username] = {"pwd": pwd}
save_user_db(db)
return "注册成功,请登录"
# 登录校验
def login_check(username, pwd):
db = load_user_db()
if username not in db:
return False, "账号不存在"
if db[username]["pwd"] != pwd:
return False, "密码错误"
return True, "登录成功"
# 获取用户文件列表
def get_user_files(username):
user_dir = f"users/{username}/"
all_files = list_repo_files(repo_id=DATASET_REPO, repo_type="dataset")
user_files = []
for f in all_files:
if f.startswith(user_dir):
file_name = f.replace(user_dir, "")
user_files.append((file_name, f))
return user_files
# 上传文件到用户独立目录
def upload_user_file(username, file_obj):
user_dir = f"users/{username}/"
file_name = os.path.basename(file_obj.name)
target_path = user_dir + file_name
upload_file(path_or_fileobj=file_obj, path_in_repo=target_path, repo_id=DATASET_REPO, repo_type="dataset")
return f"上传成功:{file_name}"
# 下载文件
def download_user_file(username, file_name):
user_dir = f"users/{username}/"
remote_path = user_dir + file_name
local_save = f"./download_{file_name}"
download_file(repo_id=DATASET_REPO, filename=remote_path, local_dir="./", repo_type="dataset")
return local_save
# 删除文件
def del_user_file(username, file_name):
user_dir = f"users/{username}/"
remote_path = user_dir + file_name
delete_file(path_in_repo=remote_path, repo_id=DATASET_REPO, repo_type="dataset")
return f"已删除:{file_name}"
# Gradio界面
with gr.Blocks(title="HF云端网盘") as demo:
gr.Markdown("# 云端网盘(账号登录+持久文件存储)")
username = gr.Textbox(label="用户名")
pwd = gr.Textbox(label="密码", type="password")
with gr.Row():
login_btn = gr.Button("登录")
reg_btn = gr.Button("注册")
msg = gr.Textbox(label="提示信息", interactive=False)
# 网盘区域(登录后显示)
with gr.Column(visible=False) as disk_panel:
gr.Markdown("## 文件管理")
upload_input = gr.File(label="上传文件", file_count="multiple")
upload_msg = gr.Textbox(label="上传结果", interactive=False)
file_drop = gr.Dropdown(label="已有文件", choices=[])
with gr.Row():
dl_btn = gr.Button("下载选中文件")
del_btn = gr.Button("删除选中文件")
dl_out = gr.File(label="下载文件")
# 注册事件
def reg_action(u, p):
res = register(u, p)
return res
reg_btn.click(reg_action, inputs=[username, pwd], outputs=[msg])
# 登录事件
def login_action(u, p):
ok, tip = login_check(u, p)
if ok:
files = get_user_files(u)
names = [i[0] for i in files]
return tip, gr.update(visible=True), gr.update(choices=names)
else:
return tip, gr.update(visible=False), gr.update(choices=[])
login_btn.click(login_action, inputs=[username, pwd], outputs=[msg, disk_panel, file_drop])
# 上传
def upload_action(u, files):
tips = []
for f in files:
t = upload_user_file(u, f)
tips.append(t)
new_files = get_user_files(u)
names = [i[0] for i in new_files]
return "\n".join(tips), gr.update(choices=names)
upload_input.change(upload_action, inputs=[username, upload_input], outputs=[upload_msg, file_drop])
# 下载
def dl_action(u, fname):
path = download_user_file(u, fname)
return path
dl_btn.click(dl_action, inputs=[username, file_drop], outputs=[dl_out])
# 删除
def del_action(u, fname):
tip = del_user_file(u, fname)
new_files = get_user_files(u)
names = [i[0] for i in new_files]
return tip, gr.update(choices=names)
del_btn.click(del_action, inputs=[username, file_drop], outputs=[upload_msg, file_drop])
if __name__ == "__main__":
demo.launch(server_name="0.0.0.0")