Spaces:
Sleeping
Sleeping
File size: 5,972 Bytes
e989cbd 167589c f03947f 5504b0a f03947f 5504b0a f03947f 5504b0a f03947f 5504b0a e989cbd | 1 2 3 4 5 6 7 8 9 10 11 12 13 14 15 16 17 18 19 20 21 22 23 24 25 26 27 28 29 30 31 32 33 34 35 36 37 38 39 40 41 42 43 44 45 46 47 48 49 50 51 52 53 54 55 56 57 58 59 60 61 62 63 64 65 66 67 68 69 70 71 72 73 74 75 76 77 78 79 80 81 82 83 84 85 86 87 88 89 90 91 92 93 94 95 96 97 98 99 100 101 102 103 104 105 106 107 108 109 110 111 112 113 114 115 116 117 118 119 120 121 122 123 124 125 126 127 128 129 130 131 132 133 134 135 136 137 138 139 140 141 142 143 144 145 146 147 148 149 150 151 152 153 154 155 156 157 158 159 160 161 162 163 164 165 166 167 168 169 170 171 172 173 174 175 | import sqlite3
import json
import threading
from pathlib import Path
import time
DB_PATH = Path("database.db")
db_lock = threading.Lock()
class DatabaseManager:
def __init__(self):
self.init_db()
def get_connection(self):
conn = sqlite3.connect(DB_PATH, check_same_thread=False)
conn.row_factory = sqlite3.Row
return conn
def init_db(self):
with db_lock:
conn = self.get_connection()
cursor = conn.cursor()
# Tasks Table
cursor.execute("""
CREATE TABLE IF NOT EXISTS tasks (
task_id TEXT PRIMARY KEY,
user_email TEXT,
status TEXT,
progress INTEGER,
step TEXT,
message TEXT,
created_at REAL,
start_time REAL,
completed_time REAL,
queue_position INTEGER,
eta_seconds REAL,
result_path TEXT,
error TEXT,
metadata TEXT
)
""")
# Users Table (for quotas, premium status)
cursor.execute("""
CREATE TABLE IF NOT EXISTS users (
email TEXT PRIMARY KEY,
is_premium BOOLEAN DEFAULT 0,
credits INTEGER DEFAULT 10,
last_reset REAL
)
""")
conn.commit()
conn.close()
def get_task(self, task_id: str):
with db_lock:
conn = self.get_connection()
cursor = conn.cursor()
cursor.execute("SELECT * FROM tasks WHERE task_id = ?", (task_id,))
row = cursor.fetchone()
conn.close()
if row:
d = dict(row)
if d.get("metadata"):
try:
d["metadata"] = json.loads(d["metadata"])
except:
pass
return d
return None
def upsert_task(self, task_id: str, data: dict):
with db_lock:
conn = self.get_connection()
cursor = conn.cursor()
cursor.execute("SELECT task_id FROM tasks WHERE task_id = ?", (task_id,))
exists = cursor.fetchone()
metadata_str = json.dumps(data.get("metadata", {})) if "metadata" in data else None
if exists:
updates = []
values = []
for k, v in data.items():
if k == "task_id" or k == "metadata": continue
updates.append(f"{k} = ?")
values.append(v)
if "metadata" in data:
updates.append("metadata = ?")
values.append(metadata_str)
if updates:
values.append(task_id)
query = f"UPDATE tasks SET {', '.join(updates)} WHERE task_id = ?"
cursor.execute(query, values)
else:
columns = []
values = []
placeholders = []
for k, v in data.items():
if k == "metadata": continue
columns.append(k)
values.append(v)
placeholders.append("?")
if "metadata" in data:
columns.append("metadata")
values.append(metadata_str)
placeholders.append("?")
columns.append("task_id")
values.append(task_id)
placeholders.append("?")
query = f"INSERT INTO tasks ({', '.join(columns)}) VALUES ({', '.join(placeholders)})"
cursor.execute(query, values)
conn.commit()
conn.close()
def get_user(self, email: str):
with db_lock:
conn = self.get_connection()
cursor = conn.cursor()
cursor.execute("SELECT * FROM users WHERE email = ?", (email,))
row = cursor.fetchone()
conn.close()
return dict(row) if row else None
def ensure_user(self, email: str):
user = self.get_user(email)
if not user:
with db_lock:
conn = self.get_connection()
cursor = conn.cursor()
cursor.execute("INSERT INTO users (email, is_premium, credits, last_reset) VALUES (?, 0, 10, ?)", (email, time.time()))
conn.commit()
conn.close()
return self.get_user(email)
return user
def fail_stuck_tasks(self):
"""Marks any tasks stuck in 'queued' or 'processing' as failed when server restarts"""
with db_lock:
conn = self.get_connection()
cursor = conn.cursor()
cursor.execute(
"UPDATE tasks SET status = 'failed', message = 'Server restarted. Please try uploading again.' WHERE status IN ('queued', 'processing')"
)
conn.commit()
conn.close()
def get_all_tasks_for_user(self, email: str):
with db_lock:
conn = self.get_connection()
cursor = conn.cursor()
cursor.execute("SELECT * FROM tasks WHERE user_email = ? ORDER BY created_at DESC", (email,))
rows = cursor.fetchall()
conn.close()
tasks = []
for row in rows:
d = dict(row)
if d.get("metadata"):
try:
d["metadata"] = json.loads(d["metadata"])
except:
pass
tasks.append(d)
return tasks
db_manager = DatabaseManager()
|