Me / db_helper.py
FrnklnWrld's picture
Update db_helper.py
7b56b6c verified
Raw
History Blame Contribute Delete
12.9 kB
# db_helper.py
# Upload this file to your HuggingFace Space to replace the existing one.
# Changes from original:
# - get_user_messages: fixed field name (message_type, not is_from_bot)
# - Added: create_pool(), get_pool_by_token() for Engine 3 secondary writes
# - Added: increment_pool_response_count() for tracking
# - Fixed: store self.url directly instead of relying on
# db.client.options.url, which doesn't exist on SyncClientOptions
# in the currently installed supabase-py version
# Install: pip install supabase python-dotenv
import os
from supabase import create_client, Client
from dotenv import load_dotenv
from datetime import datetime, timedelta
from typing import Optional, List, Dict
load_dotenv()
class DB:
"""Simple helper class for Supabase operations"""
def __init__(self):
load_dotenv()
url = os.getenv("SUPABASE_URL", "https://jjnbusmjsgjyjgomhuij.supabase.co")
key = os.getenv("SUPABASE_SERVICE_KEY")
if not key:
raise ValueError("SUPABASE_SERVICE_KEY not set in HF Secrets")
self.url = url
self.client: Client = create_client(url, key)
# ==================== USERS ====================
def create_user(self, display_name: str, email: Optional[str] = None) -> dict:
"""Create a new user"""
data = {"display_name": display_name, "email": email}
result = self.client.table("users").insert(data).execute()
return result.data[0] if result.data else None
def get_user(self, display_name: str) -> Optional[dict]:
"""Get user by display name"""
result = (
self.client.table("users")
.select("*")
.eq("display_name", display_name)
.execute()
)
return result.data[0] if result.data else None
def get_or_create_user(self, display_name: str, email: Optional[str] = None) -> dict:
"""Get existing user or create new one"""
user = self.get_user(display_name)
if user:
return user
return self.create_user(display_name, email)
def update_user_last_seen(self, display_name: str) -> dict:
"""Update user's last seen timestamp"""
result = (
self.client.table("users")
.update({"last_seen": datetime.now().isoformat()})
.eq("display_name", display_name)
.execute()
)
return result.data[0] if result.data else None
def get_all_users(self) -> List[dict]:
"""Get all users"""
result = (
self.client.table("users")
.select("*")
.order("created_at", desc=True)
.execute()
)
return result.data
def get_user_messages(self, user_id: str, limit: int = 5) -> List[Dict]:
"""Get recent chat messages for a user (used in /journey endpoint)"""
try:
result = (
self.client.table("chat_history")
.select("*")
.eq("user_id", user_id)
.order("happened_at", desc=True)
.limit(limit)
.execute()
)
return result.data
except Exception as e:
print(f"Error fetching user messages: {str(e)}")
return []
# ==================== JOURNEYS ====================
def create_journey(self, user_id: str, category: str) -> dict:
"""Create a new journey for a user in a category"""
data = {"user_id": user_id, "category": category, "status": "active"}
result = self.client.table("journeys").insert(data).execute()
return result.data[0] if result.data else None
def get_journey(self, user_id: str, category: str) -> Optional[dict]:
"""Get a specific journey by user and category"""
result = (
self.client.table("journeys")
.select("*")
.eq("user_id", user_id)
.eq("category", category)
.execute()
)
return result.data[0] if result.data else None
def get_or_create_journey(self, user_id: str, category: str) -> dict:
"""Get existing journey or create new one"""
journey = self.get_journey(user_id, category)
if journey:
return journey
return self.create_journey(user_id, category)
def get_user_journeys(self, user_id: str) -> List[dict]:
"""Get all journeys for a user"""
result = (
self.client.table("journeys")
.select("*")
.eq("user_id", user_id)
.order("last_active_at", desc=True)
.execute()
)
return result.data
def update_journey(self, journey_id: str, **kwargs) -> dict:
"""Update journey fields"""
kwargs["last_active_at"] = datetime.now().isoformat()
result = (
self.client.table("journeys")
.update(kwargs)
.eq("id", journey_id)
.execute()
)
return result.data[0] if result.data else None
def update_pending_questions(self, journey_id: str, pending_list: List[Dict]):
"""Save pending MCQs as jsonb array"""
self.client.table("journeys").update(
{"pending_questions": pending_list}
).eq("id", journey_id).execute()
def get_pending_questions(self, journey_id: str) -> List[Dict]:
"""Load pending MCQs from jsonb"""
result = (
self.client.table("journeys")
.select("pending_questions")
.eq("id", journey_id)
.execute()
)
if result.data and result.data[0].get("pending_questions"):
return result.data[0]["pending_questions"]
return []
# ==================== MCQ ANSWERS ====================
def add_answer(
self,
journey_id: str,
question: str,
answer: str,
score: int,
batch_number: Optional[int] = None,
) -> dict:
"""Add a new MCQ answer"""
data = {
"journey_id": journey_id,
"question_text": question,
"answer_text": answer,
"score": score,
"batch_number": batch_number,
}
result = self.client.table("mcq_answers").insert(data).execute()
return result.data[0] if result.data else None
def get_journey_answers(self, journey_id: str, limit: Optional[int] = None) -> List[dict]:
"""Get all answers for a journey"""
query = (
self.client.table("mcq_answers")
.select("*")
.eq("journey_id", journey_id)
.order("answered_at", desc=True)
)
if limit:
query = query.limit(limit)
result = query.execute()
return result.data
def get_journey_stats(self, journey_id: str) -> dict:
"""Get statistics for a journey"""
answers = self.get_journey_answers(journey_id)
if not answers:
return {"total_answers": 0, "avg_score": 0}
total = len(answers)
avg_score = sum(a["score"] for a in answers) / total
return {
"total_answers": total,
"avg_score": round(avg_score, 2),
"latest_score": answers[0]["score"],
"latest_answer_time": answers[0]["answered_at"],
}
def add_batch_answers(self, journey_id: str, answers_data: List[Dict]) -> List[dict]:
"""Add multiple answers at once"""
batch_number = int(datetime.now().timestamp())
data = [
{
"journey_id": journey_id,
"question_text": item["question"],
"answer_text": item["answer"],
"score": item["score"],
"batch_number": batch_number,
}
for item in answers_data
]
result = self.client.table("mcq_answers").insert(data).execute()
return result.data
# ==================== CHAT HISTORY ====================
def add_chat_message(
self,
user_id: str,
content: str,
is_from_bot: bool = False,
category_context: Optional[str] = None,
) -> dict:
"""Add a chat message"""
data = {
"user_id": user_id,
"message_type": "bot" if is_from_bot else "user",
"content": content,
"category_context": category_context,
}
result = self.client.table("chat_history").insert(data).execute()
return result.data[0] if result.data else None
def get_chat_history(
self, user_id: str, limit: int = 50, category: Optional[str] = None
) -> List[dict]:
"""Get chat history for a user"""
query = (
self.client.table("chat_history").select("*").eq("user_id", user_id)
)
if category:
query = query.eq("category_context", category)
result = query.order("happened_at", desc=False).limit(limit).execute()
return result.data
# ==================== REMINDERS ====================
def create_reminder(
self,
user_id: str,
scheduled_for: datetime,
reminder_type: str,
journey_id: Optional[str] = None,
) -> dict:
"""Create a new reminder"""
data = {
"user_id": user_id,
"journey_id": journey_id,
"scheduled_for": scheduled_for.isoformat(),
"type": reminder_type,
"status": "pending",
}
result = self.client.table("reminders").insert(data).execute()
return result.data[0] if result.data else None
def get_pending_reminders(self, user_id: Optional[str] = None) -> List[dict]:
"""Get all pending reminders, optionally for a specific user"""
query = (
self.client.table("reminders")
.select("*")
.eq("status", "pending")
.lte("scheduled_for", datetime.now().isoformat())
)
if user_id:
query = query.eq("user_id", user_id)
result = query.order("scheduled_for").execute()
return result.data
def get_user_reminders(self, user_id: str) -> List[dict]:
"""Get all reminders for a user"""
result = (
self.client.table("reminders")
.select("*")
.eq("user_id", user_id)
.order("scheduled_for", desc=True)
.execute()
)
return result.data
def update_reminder_status(self, reminder_id: str, status: str) -> dict:
"""Update reminder status"""
result = (
self.client.table("reminders")
.update({"status": status})
.eq("id", reminder_id)
.execute()
)
return result.data[0] if result.data else None
# ==================== POOLS (Engine 3 secondary write) ====================
def create_pool(
self,
owner_id: str,
category: str,
share_token: str,
question_ids: List[str],
min_threshold: int = 5,
) -> Optional[dict]:
"""
Create a Muhasabah Pool entry in Supabase.
Primary store is Firestore — this is a secondary copy for analytics.
Returns None silently if the pools table doesn't exist yet.
"""
try:
data = {
"owner_id": owner_id,
"category": category,
"share_token": share_token,
"question_ids": question_ids,
"min_threshold": min_threshold,
}
result = self.client.table("pools").insert(data).execute()
return result.data[0] if result.data else None
except Exception as e:
print(f"[pools] Non-fatal: could not write pool to Supabase: {e}")
return None
def get_pool_by_token(self, share_token: str) -> Optional[dict]:
"""Get pool by share token"""
try:
result = (
self.client.table("pools")
.select("*")
.eq("share_token", share_token)
.execute()
)
return result.data[0] if result.data else None
except Exception as e:
print(f"[pools] Non-fatal: {e}")
return None
def increment_pool_response_count(self, share_token: str) -> None:
"""Increment the response count for a pool after a submission"""
try:
pool = self.get_pool_by_token(share_token)
if pool:
new_count = (pool.get("response_count") or 0) + 1
self.client.table("pools").update(
{"response_count": new_count}
).eq("share_token", share_token).execute()
except Exception as e:
print(f"[pools] Non-fatal: {e}")