From cfeb67426ae2abc3636beb205dc02d0a5dfb2cbf Mon Sep 17 00:00:00 2001 From: Pierre Brunelle Date: Sat, 12 Sep 2026 15:28:09 -0700 Subject: [PATCH] Simplify API error and notification paths --- AGENTS.md | 3 + README.md | 1 + backend/pixelbot/functions.py | 42 +- backend/pixelbot/models.py | 12 - backend/pixelbot/notifications.py | 86 ++ backend/pixelbot/routers/chat.py | 242 ++-- backend/pixelbot/routers/database.py | 663 +++++----- backend/pixelbot/routers/experiments.py | 116 +- backend/pixelbot/routers/export.py | 43 +- backend/pixelbot/routers/files.py | 440 +++---- backend/pixelbot/routers/history.py | 354 ++---- backend/pixelbot/routers/images.py | 806 +++++------- backend/pixelbot/routers/integrations.py | 77 +- backend/pixelbot/routers/memory.py | 23 +- backend/pixelbot/routers/personas.py | 108 +- backend/pixelbot/routers/studio.py | 1456 ++++++++++------------ backend/pixelbot/utils.py | 75 +- backend/pyproject.toml | 1 + backend/tests/conftest.py | 23 +- backend/tests/test_contract.py | 50 +- backend/tests/test_schema_cli.py | 3 +- docs/pixeltable-0.7.7-upgrade.md | 6 + frontend/src/types/index.ts | 24 - 23 files changed, 2003 insertions(+), 2651 deletions(-) create mode 100644 backend/pixelbot/notifications.py diff --git a/AGENTS.md b/AGENTS.md index 561b633..55b2db9 100644 --- a/AGENTS.md +++ b/AGENTS.md @@ -7,6 +7,7 @@ Pixelbot 3.0 is an app-first Pixeltable 0.7.7 project. Python 3.11+ and Node 22. - `backend/pixelbot/app.py`: application entry, endpoint-facing `@pxt.query` functions, module-scope `FastAPIRouter` routes, and the custom FastAPI app. - `backend/pixelbot/schema.py`: the single `TableModel` schema for `pixelbot_v3`; its query functions are reusable parts of computed pipelines. - `backend/pixelbot/routers/`: custom REST writes and HTTP-specific behavior. The database router is read-only. +- `backend/pixelbot/notifications.py`: the single notification transport used by both HTTP routes and Pixeltable tools. - `backend/pixelbot/static/`: production SPA package data generated by Vite. - `frontend/`: React SPA with route-level lazy loading. - Pixeltable agent guidance lives in the canonical @@ -21,6 +22,8 @@ Declare stored columns with annotations and computed columns with assignments in Use `@pxt.query` plus `FastAPIRouter.add_query_route()` for table-backed read endpoints. Put endpoint-only queries beside their routes in `app.py`; keep queries in `schema.py` when computed columns also call them. Use custom FastAPI handlers for request validation, application defaults, and multi-step writes. Do not mirror a query in a plain Python helper. +Let the application exception handler log and sanitize unexpected route failures. Convert only expected validation and missing-resource failures to `HTTPException`. Apply `pxt_retry()` only to read-only operations; retrying writes, provider calls, or notifications can duplicate side effects. + A changed computed expression is unsupported by schema migration. Rename the column, or drop and re-add it in separate schema updates. `--allow-destructive` does not change that restriction. ## Workflow diff --git a/README.md b/README.md index 42525b7..1d247a4 100644 --- a/README.md +++ b/README.md @@ -64,6 +64,7 @@ The Database page is an inspector for catalog rows, schemas, lineage, history, s - Media reads are confined to the upload root and configured `PIXELTABLE_HOME`, including resolved symlinks. - Database inspection and exports are confined to `pixelbot_v3` and registered scratch tables, with result caps. - The agent webhook tool can send only to `WEBHOOK_URL` configured at process startup. +- HTTP routes and agent tools share one notification transport, and catalog logs store only the destination origin. - There is no runtime expression evaluation or general catalog mutation API. ## Validation diff --git a/backend/pixelbot/functions.py b/backend/pixelbot/functions.py index 5fd7490..20fb12a 100644 --- a/backend/pixelbot/functions.py +++ b/backend/pixelbot/functions.py @@ -9,6 +9,8 @@ import yfinance as yf from duckduckgo_search import DDGS +from pixelbot.notifications import deliver_notification + @pxt.udf def get_latest_news(topic: str) -> str: @@ -186,47 +188,25 @@ def fetch_financial_data(ticker: str) -> str: @pxt.udf def send_slack_message(message: str) -> str: """Send a message to a configured Slack channel via incoming webhook.""" - webhook_url = os.environ.get("SLACK_WEBHOOK_URL", "") - if not webhook_url: - return "Error: SLACK_WEBHOOK_URL not configured." - try: - resp = requests.post(webhook_url, json={"text": message}, timeout=10) - if resp.status_code == 200: - return "Slack message sent successfully." - return f"Slack error ({resp.status_code}): {resp.text}" - except requests.RequestException as e: - return f"Slack request failed: {e}" + result = deliver_notification("slack", message) + assert result is not None + return result.message @pxt.udf def send_discord_message(message: str) -> str: """Send a message to a configured Discord channel via webhook.""" - webhook_url = os.environ.get("DISCORD_WEBHOOK_URL", "") - if not webhook_url: - return "Error: DISCORD_WEBHOOK_URL not configured." - try: - resp = requests.post(webhook_url, json={"content": message}, timeout=10) - if resp.status_code in (200, 204): - return "Discord message sent successfully." - return f"Discord error ({resp.status_code}): {resp.text}" - except requests.RequestException as e: - return f"Discord request failed: {e}" + result = deliver_notification("discord", message) + assert result is not None + return result.message @pxt.udf def send_webhook(message: str) -> str: """POST a JSON payload to the configured webhook URL.""" - target_url = os.environ.get("WEBHOOK_URL", "") - if not target_url: - return "Error: WEBHOOK_URL not configured." - try: - payload = {"text": message, "source": "pixelbot", "timestamp": datetime.utcnow().isoformat()} - resp = requests.post(target_url, json=payload, timeout=10) - if resp.status_code < 300: - return f"Webhook delivered ({resp.status_code})." - return f"Webhook error ({resp.status_code}): {resp.text}" - except requests.RequestException as e: - return f"Webhook request failed: {e}" + result = deliver_notification("webhook", message) + assert result is not None + return result.message @pxt.udf diff --git a/backend/pixelbot/models.py b/backend/pixelbot/models.py index d0f1721..d23decb 100644 --- a/backend/pixelbot/models.py +++ b/backend/pixelbot/models.py @@ -263,18 +263,6 @@ class AddUrlResponse(BaseModel): uuid: str -class DeleteFileResponse(BaseModel): - message: str - db_deleted: bool - file_deleted: bool - uuid: str - - -class DeleteAllResponse(BaseModel): - message: str - should_refresh: bool = True - - # ── History ────────────────────────────────────────────────────────────────── diff --git a/backend/pixelbot/notifications.py b/backend/pixelbot/notifications.py new file mode 100644 index 0000000..efe57a6 --- /dev/null +++ b/backend/pixelbot/notifications.py @@ -0,0 +1,86 @@ +"""Notification delivery shared by the HTTP API and Pixeltable agent tools.""" + +from __future__ import annotations + +import logging +from dataclasses import dataclass +from datetime import UTC, datetime +from urllib.parse import urlsplit + +import requests + +from pixelbot import config + +logger = logging.getLogger(__name__) + + +@dataclass(frozen=True) +class DeliveryResult: + message: str + success: bool + response_code: int + + +def deliver_notification(service: str, message: str) -> DeliveryResult | None: + """Deliver a message to one configured notification service.""" + service = service.lower() + url = _service_url(service) + if service not in config.INTEGRATIONS: + return None + if not url: + env_var = config.INTEGRATIONS[service]["env_var"] + return DeliveryResult(f"Error: {env_var} not configured.", False, 0) + + payload = _payload(service, message) + try: + response = requests.post(url, json=payload, timeout=10) + except requests.RequestException: + logger.exception("%s notification request failed", service) + return DeliveryResult(f"{service.title()} request failed.", False, 0) + + success = response.status_code in _success_codes(service) + if success: + if service == "webhook": + result_message = f"Webhook delivered ({response.status_code})." + else: + result_message = f"{service.title()} message sent successfully." + else: + result_message = f"{service.title()} delivery failed ({response.status_code})." + return DeliveryResult(result_message, success, response.status_code) + + +def redacted_destination(service: str) -> str: + """Return a display-safe form of the configured destination.""" + url = _service_url(service.lower()) + if not url: + return "(not configured)" + parsed = urlsplit(url) + return f"{parsed.scheme}://{parsed.netloc}/..." + + +def _service_url(service: str) -> str: + return { + "slack": config.SLACK_WEBHOOK_URL, + "discord": config.DISCORD_WEBHOOK_URL, + "webhook": config.WEBHOOK_URL, + }.get(service, "") + + +def _payload(service: str, message: str) -> dict[str, str]: + if service == "discord": + return {"content": message} + if service == "webhook": + return { + "text": message, + "source": "pixelbot", + "timestamp": datetime.now(UTC).isoformat(), + } + return {"text": message} + + +def _success_codes(service: str) -> range | tuple[int, ...]: + if service == "slack": + return (200,) + if service == "discord": + return (200, 204) + return range(200, 300) diff --git a/backend/pixelbot/routers/chat.py b/backend/pixelbot/routers/chat.py index bef454e..ff2ee20 100644 --- a/backend/pixelbot/routers/chat.py +++ b/backend/pixelbot/routers/chat.py @@ -7,7 +7,6 @@ from pixelbot import config from pixelbot.models import ChatHistoryRow, QueryMetadata, QueryResponse, ToolAgentRow -from pixelbot.utils import pxt_retry logger = logging.getLogger(__name__) router = APIRouter(prefix="/api", tags=["chat"]) @@ -20,146 +19,137 @@ class QueryRequest(BaseModel): @router.post("/query", response_model=QueryResponse) -@pxt_retry() def query(body: QueryRequest): """Process a user query through the Pixeltable agent workflow.""" user_id = config.DEFAULT_USER_ID if not body.query: raise HTTPException(status_code=400, detail="Query text is required") + tool_agent = pxt.get_table("pixelbot_v3.tools") - try: - tool_agent = pxt.get_table("pixelbot_v3.tools") - - # Determine prompts and parameters - selected_initial_prompt = config.INITIAL_SYSTEM_PROMPT - selected_final_prompt = config.FINAL_SYSTEM_PROMPT - selected_max_tokens = config.DEFAULT_MAX_TOKENS - selected_temperature = config.DEFAULT_TEMPERATURE - - if body.persona_id: - try: - personas_table = pxt.get_table("pixelbot_v3.user_personas") - persona_result = ( - personas_table.where( - (personas_table.user_id == user_id) & (personas_table.persona_name == body.persona_id) - ) - .select( - initial_prompt=personas_table.initial_prompt, - final_prompt=personas_table.final_prompt, - llm_params=personas_table.llm_params, + # Determine prompts and parameters + selected_initial_prompt = config.INITIAL_SYSTEM_PROMPT + selected_final_prompt = config.FINAL_SYSTEM_PROMPT + selected_max_tokens = config.DEFAULT_MAX_TOKENS + selected_temperature = config.DEFAULT_TEMPERATURE + + if body.persona_id: + try: + personas_table = pxt.get_table("pixelbot_v3.user_personas") + persona_result = ( + personas_table.where( + (personas_table.user_id == user_id) & (personas_table.persona_name == body.persona_id) + ) + .select( + initial_prompt=personas_table.initial_prompt, + final_prompt=personas_table.final_prompt, + llm_params=personas_table.llm_params, + ) + .collect() + ) + if len(persona_result) > 0: + custom_data = persona_result[0] + selected_initial_prompt = custom_data["initial_prompt"] + selected_final_prompt = custom_data["final_prompt"] + llm_params = custom_data.get("llm_params") or {} + selected_max_tokens = llm_params.get("max_tokens", selected_max_tokens) + selected_temperature = llm_params.get("temperature", selected_temperature) + logger.info(f"Loaded persona '{body.persona_id}' for user {user_id}") + except Exception as db_err: + logger.error(f"Error fetching persona '{body.persona_id}': {db_err}", exc_info=True) + + # Insert the query using a validated Pydantic model + current_timestamp = datetime.now() + row = ToolAgentRow( + prompt=body.query, + timestamp=current_timestamp, + user_id=user_id, + initial_system_prompt=selected_initial_prompt, + final_system_prompt=selected_final_prompt, + max_tokens=selected_max_tokens, + temperature=selected_temperature, + ) + # insert() blocks until all computed columns finish; return_rows avoids follow-up query + status = tool_agent.insert([row], return_rows=True) + + if not status.rows: + raise HTTPException(status_code=500, detail="No results found after processing query") + + result_data = status.rows[0] + + # Process image context + processed_image_context: list[dict] = [] + if result_data.get("image_context"): + for item in result_data["image_context"]: + if isinstance(item, dict) and "encoded_image" in item and item["encoded_image"]: + encoded = item["encoded_image"] + if isinstance(encoded, bytes): + encoded = encoded.decode("utf-8") + if isinstance(encoded, str) and encoded: + processed_image_context.append({"encoded_image": encoded}) + + # Process video frame context + processed_video_frame_context: list[dict] = [] + if result_data.get("video_frame_context"): + for item in result_data["video_frame_context"]: + if isinstance(item, dict) and "encoded_frame" in item and item["encoded_frame"]: + frame_data = item["encoded_frame"] + if isinstance(frame_data, bytes): + frame_data = frame_data.decode("utf-8") + if isinstance(frame_data, str) and frame_data: + processed_video_frame_context.append( + { + "encoded_frame": frame_data, + "sim": item.get("sim"), + "timestamp": item.get("timestamp"), + } ) - .collect() + + # Insert into chat history using validated Pydantic models + conversation_id = body.conversation_id or "default" + try: + chat_history_table = pxt.get_table("pixelbot_v3.chat_history") + chat_history_table.insert( + [ + ChatHistoryRow( + role="user", + content=body.query, + conversation_id=conversation_id, + timestamp=current_timestamp, + user_id=user_id, ) - if len(persona_result) > 0: - custom_data = persona_result[0] - selected_initial_prompt = custom_data["initial_prompt"] - selected_final_prompt = custom_data["final_prompt"] - llm_params = custom_data.get("llm_params") or {} - selected_max_tokens = llm_params.get("max_tokens", selected_max_tokens) - selected_temperature = llm_params.get("temperature", selected_temperature) - logger.info(f"Loaded persona '{body.persona_id}' for user {user_id}") - except Exception as db_err: - logger.error(f"Error fetching persona '{body.persona_id}': {db_err}", exc_info=True) - - # Insert the query using a validated Pydantic model - current_timestamp = datetime.now() - row = ToolAgentRow( - prompt=body.query, - timestamp=current_timestamp, - user_id=user_id, - initial_system_prompt=selected_initial_prompt, - final_system_prompt=selected_final_prompt, - max_tokens=selected_max_tokens, - temperature=selected_temperature, + ] ) - # insert() blocks until all computed columns finish; return_rows avoids follow-up query - status = tool_agent.insert([row], return_rows=True) - - if not status.rows: - raise HTTPException(status_code=500, detail="No results found after processing query") - - result_data = status.rows[0] - - # Process image context - processed_image_context: list[dict] = [] - if result_data.get("image_context"): - for item in result_data["image_context"]: - if isinstance(item, dict) and "encoded_image" in item and item["encoded_image"]: - encoded = item["encoded_image"] - if isinstance(encoded, bytes): - encoded = encoded.decode("utf-8") - if isinstance(encoded, str) and encoded: - processed_image_context.append({"encoded_image": encoded}) - - # Process video frame context - processed_video_frame_context: list[dict] = [] - if result_data.get("video_frame_context"): - for item in result_data["video_frame_context"]: - if isinstance(item, dict) and "encoded_frame" in item and item["encoded_frame"]: - frame_data = item["encoded_frame"] - if isinstance(frame_data, bytes): - frame_data = frame_data.decode("utf-8") - if isinstance(frame_data, str) and frame_data: - processed_video_frame_context.append( - { - "encoded_frame": frame_data, - "sim": item.get("sim"), - "timestamp": item.get("timestamp"), - } - ) - - # Insert into chat history using validated Pydantic models - conversation_id = body.conversation_id or "default" - try: - chat_history_table = pxt.get_table("pixelbot_v3.chat_history") + answer = result_data.get("answer", "Error: Answer not generated.") + if answer and not answer.startswith("Error:"): chat_history_table.insert( [ ChatHistoryRow( - role="user", - content=body.query, + role="assistant", + content=answer, conversation_id=conversation_id, - timestamp=current_timestamp, + timestamp=datetime.now(), user_id=user_id, ) ] ) - answer = result_data.get("answer", "Error: Answer not generated.") - if answer and not answer.startswith("Error:"): - chat_history_table.insert( - [ - ChatHistoryRow( - role="assistant", - content=answer, - conversation_id=conversation_id, - timestamp=datetime.now(), - user_id=user_id, - ) - ] - ) - except Exception as history_err: - logger.error(f"Error inserting into chat history: {history_err}") - - metadata = QueryMetadata( - timestamp=current_timestamp.isoformat(), - has_doc_context=bool(result_data.get("doc_context")), - has_image_context=bool(result_data.get("image_context")), - has_tool_output=bool(result_data.get("tool_output")), - has_history_context=bool(result_data.get("history_context")), - has_memory_context=bool(result_data.get("memory_context")), - has_chat_memory_context=bool(result_data.get("chat_memory_context")), - ) - - return QueryResponse( - answer=result_data.get("answer", "Error: Answer not generated."), - metadata=metadata, - image_context=processed_image_context, - video_frame_context=processed_video_frame_context, - follow_up_text=result_data.get("follow_up_text"), - ) - - except HTTPException: - raise - except Exception as e: - logger.error(f"Error processing query: {e}", exc_info=True) - raise HTTPException(status_code=500, detail=str(e)) + except Exception as history_err: + logger.error(f"Error inserting into chat history: {history_err}") + + metadata = QueryMetadata( + timestamp=current_timestamp.isoformat(), + has_doc_context=bool(result_data.get("doc_context")), + has_image_context=bool(result_data.get("image_context")), + has_tool_output=bool(result_data.get("tool_output")), + has_history_context=bool(result_data.get("history_context")), + has_memory_context=bool(result_data.get("memory_context")), + has_chat_memory_context=bool(result_data.get("chat_memory_context")), + ) + + return QueryResponse( + answer=result_data.get("answer", "Error: Answer not generated."), + metadata=metadata, + image_context=processed_image_context, + video_frame_context=processed_video_frame_context, + follow_up_text=result_data.get("follow_up_text"), + ) diff --git a/backend/pixelbot/routers/database.py b/backend/pixelbot/routers/database.py index 2e31934..b8a0fde 100644 --- a/backend/pixelbot/routers/database.py +++ b/backend/pixelbot/routers/database.py @@ -64,44 +64,39 @@ def _column_info(tbl) -> list[dict]: @pxt_retry() def list_all_tables(): """List all tables and views in the agents namespace with schema info.""" - try: - table_paths = pxt.list_tables(NAMESPACE, recursive=True) - - tables = [] - for path in sorted(table_paths): - try: - tbl = pxt.get_table(path) - row_count = tbl.count() - meta = tbl.get_metadata() - base_path = meta.get("base") + table_paths = pxt.list_tables(NAMESPACE, recursive=True) - tables.append( - { - "path": path, - "type": "view" if meta.get("is_view") else "table", - "base_table": base_path, - "columns": _column_info(tbl), - "row_count": row_count, - } - ) - except Exception as e: - logger.warning(f"Could not inspect table {path}: {e}") - tables.append( - { - "path": path, - "type": "unknown", - "base_table": None, - "columns": [], - "row_count": 0, - "error": str(e), - } - ) + tables = [] + for path in sorted(table_paths): + try: + tbl = pxt.get_table(path) + row_count = tbl.count() + meta = tbl.get_metadata() + base_path = meta.get("base") - return {"namespace": NAMESPACE, "tables": tables, "count": len(tables)} + tables.append( + { + "path": path, + "type": "view" if meta.get("is_view") else "table", + "base_table": base_path, + "columns": _column_info(tbl), + "row_count": row_count, + } + ) + except Exception as e: + logger.warning(f"Could not inspect table {path}: {e}") + tables.append( + { + "path": path, + "type": "unknown", + "base_table": None, + "columns": [], + "row_count": 0, + "error": str(e), + } + ) - except Exception as e: - logger.error(f"Error listing tables: {e}", exc_info=True) - raise HTTPException(status_code=500, detail=str(e)) + return {"namespace": NAMESPACE, "tables": tables, "count": len(tables)} @router.get("/table/{path:path}/rows") @@ -115,59 +110,27 @@ def get_table_rows(path: str, limit: int = 50, offset: int = 0): tbl = pxt.get_table(path) except Exception: raise HTTPException(status_code=404, detail=f"Table '{path}' not found") + total = tbl.count() + col_names = tbl.columns() - try: - total = tbl.count() - col_names = tbl.columns() - - raw_rows = tbl.select().limit(limit, offset=offset).collect() - - rows = [] - for raw in raw_rows: - row: dict = {} - for col in col_names: - val = raw.get(col) - row[col] = _safe_value(val) - rows.append(row) - - return { - "path": path, - "columns": col_names, - "rows": rows, - "total": total, - "offset": offset, - "limit": limit, - } - - except Exception as e: - logger.error(f"Error fetching rows from {path}: {e}", exc_info=True) - raise HTTPException(status_code=500, detail=str(e)) + raw_rows = tbl.select().limit(limit, offset=offset).collect() + rows = [] + for raw in raw_rows: + row: dict = {} + for col in col_names: + val = raw.get(col) + row[col] = _safe_value(val) + rows.append(row) -@router.get("/table/{path:path}/schema") -@pxt_retry() -def get_table_schema(path: str): - """Get detailed schema for a specific table.""" - path = require_allowed_table(path) - try: - tbl = pxt.get_table(path) - except Exception: - raise HTTPException(status_code=404, detail=f"Table '{path}' not found") - - try: - meta = tbl.get_metadata() - base_path = meta.get("base") - return { - "path": path, - "type": "view" if meta.get("is_view") else "table", - "base_table": base_path, - "columns": _column_info(tbl), - "row_count": tbl.count(), - } - - except Exception as e: - logger.error(f"Error getting schema for {path}: {e}", exc_info=True) - raise HTTPException(status_code=500, detail=str(e)) + return { + "path": path, + "columns": col_names, + "rows": rows, + "total": total, + "offset": offset, + "limit": limit, + } class SampleRequest(BaseModel): @@ -194,56 +157,48 @@ def sample_table(body: SampleRequest): if not body.n and not body.fraction: raise HTTPException(status_code=400, detail="Provide either 'n' or 'fraction'") + total = tbl.count() + query = tbl.select() - try: - total = tbl.count() - query = tbl.select() - - sample_kwargs: dict = {} - if body.n is not None: - sample_kwargs["n"] = min(body.n, total) - elif body.fraction is not None: - sample_kwargs["fraction"] = max(0.0, min(1.0, body.fraction)) - - if body.seed is not None: - sample_kwargs["seed"] = body.seed - - if body.stratify_by: - col_names = tbl.columns() - if body.stratify_by not in col_names: - raise HTTPException( - status_code=400, - detail=f"Column '{body.stratify_by}' not found in {body.path}", - ) - sample_kwargs["stratify_by"] = getattr(tbl, body.stratify_by) + sample_kwargs: dict = {} + if body.n is not None: + sample_kwargs["n"] = min(body.n, total) + elif body.fraction is not None: + sample_kwargs["fraction"] = max(0.0, min(1.0, body.fraction)) - raw_rows = query.sample(**sample_kwargs).collect() + if body.seed is not None: + sample_kwargs["seed"] = body.seed + if body.stratify_by: col_names = tbl.columns() - rows = [] - for raw in raw_rows: - row = {col: _safe_value(raw.get(col)) for col in col_names} - rows.append(row) - - return { - "path": body.path, - "columns": col_names, - "rows": rows, - "sample_count": len(rows), - "total": total, - "params": { - "n": body.n, - "fraction": body.fraction, - "stratify_by": body.stratify_by, - "seed": body.seed, - }, - } - - except HTTPException: - raise - except Exception as e: - logger.error(f"Sample error for {body.path}: {e}", exc_info=True) - raise HTTPException(status_code=500, detail=str(e)) + if body.stratify_by not in col_names: + raise HTTPException( + status_code=400, + detail=f"Column '{body.stratify_by}' not found in {body.path}", + ) + sample_kwargs["stratify_by"] = getattr(tbl, body.stratify_by) + + raw_rows = query.sample(**sample_kwargs).collect() + + col_names = tbl.columns() + rows = [] + for raw in raw_rows: + row = {col: _safe_value(raw.get(col)) for col in col_names} + rows.append(row) + + return { + "path": body.path, + "columns": col_names, + "rows": rows, + "sample_count": len(rows), + "total": total, + "params": { + "n": body.n, + "fraction": body.fraction, + "stratify_by": body.stratify_by, + "seed": body.seed, + }, + } @router.get("/timeline") @@ -353,58 +308,50 @@ def join_tables(body: JoinRequest): if body.join_type not in ("inner", "left", "cross"): raise HTTPException(status_code=400, detail=f"Unsupported join type: {body.join_type}") - - try: - left_col_ref = getattr(left, body.left_column) - right_col_ref = getattr(right, body.right_column) - - # Build join - join_type = cast(Literal["inner", "left", "cross"], body.join_type) - if join_type == "cross": - joined = left.join(right, how="cross") - else: - joined = left.join(right, on=left_col_ref == right_col_ref, how=join_type) - - # Select all columns from both tables (prefix to avoid collisions) - select_kwargs = {} - for col in left_cols: - key = f"l_{col}" - try: - select_kwargs[key] = getattr(left, col) - except Exception: - pass - for col in right_cols: - key = f"r_{col}" - try: - select_kwargs[key] = getattr(right, col) - except Exception: - pass - - raw_rows = joined.select(**select_kwargs).limit(body.limit).collect() - - rows = [] - for raw in raw_rows: - row = {} - for k, v in raw.items(): - row[k] = _safe_value(v) - rows.append(row) - - return { - "left_table": body.left_table, - "right_table": body.right_table, - "join_type": body.join_type, - "left_column": body.left_column, - "right_column": body.right_column, - "columns": list(select_kwargs.keys()), - "rows": rows, - "count": len(rows), - } - - except HTTPException: - raise - except Exception as e: - logger.error(f"Join error: {e}", exc_info=True) - raise HTTPException(status_code=500, detail=str(e)) + left_col_ref = getattr(left, body.left_column) + right_col_ref = getattr(right, body.right_column) + + # Build join + join_type = cast(Literal["inner", "left", "cross"], body.join_type) + if join_type == "cross": + joined = left.join(right, how="cross") + else: + joined = left.join(right, on=left_col_ref == right_col_ref, how=join_type) + + # Select all columns from both tables (prefix to avoid collisions) + select_kwargs = {} + for col in left_cols: + key = f"l_{col}" + try: + select_kwargs[key] = getattr(left, col) + except Exception: + pass + for col in right_cols: + key = f"r_{col}" + try: + select_kwargs[key] = getattr(right, col) + except Exception: + pass + + raw_rows = joined.select(**select_kwargs).limit(body.limit).collect() + + rows = [] + for raw in raw_rows: + row = {} + for k, v in raw.items(): + row[k] = _safe_value(v) + rows.append(row) + + return { + "left_table": body.left_table, + "right_table": body.right_table, + "join_type": body.join_type, + "left_column": body.left_column, + "right_column": body.right_column, + "columns": list(select_kwargs.keys()), + "rows": rows, + "count": len(rows), + } # ── Pipeline Inspector ──────────────────────────────────────────────────────── @@ -516,220 +463,174 @@ def get_pipeline(): Includes tables, views, computed column lineage, embedding indices, version history, and per-column error counts. """ - try: - table_paths = sorted(pxt.list_tables(NAMESPACE, recursive=True)) - - nodes: list[dict] = [] - edges: list[dict] = [] - - for path in table_paths: - try: - tbl = pxt.get_table(path) - md = tbl.get_metadata() - col_meta = md.get("columns", {}) - row_count = tbl.count() - - all_col_names = set(col_meta.keys()) - - columns = [] - computed_cols = [] - insertable_cols: set[str] = set() - - short_name = path.rsplit("/", 1)[-1] - - for col_name, info in col_meta.items(): - cw = info.get("computed_with") - is_computed = cw is not None - if not is_computed: - insertable_cols.add(col_name) - defined_in = info.get("defined_in") - - cw_str = str(cw)[:200] if cw else None - func_name = _extract_func_name(cw_str) if is_computed else None - func_type = _classify_func(func_name) if func_name else None - - col_entry: dict = { - "name": col_name, - "type": info.get("type_", "unknown"), - "is_computed": is_computed, - "computed_with": cw_str, - "defined_in": defined_in, - "defined_in_self": defined_in == short_name, - "func_name": func_name, - "func_type": func_type, - } - comment = info.get("comment") - if comment: - col_entry["comment"] = comment - custom_meta = info.get("custom_metadata") - if custom_meta: - col_entry["custom_metadata"] = custom_meta - columns.append(col_entry) - if is_computed: - computed_cols.append(col_name) - - # Compute error counts for computed columns (sample first 500 rows) - total_errors = 0 - for col in columns: - if col["is_computed"]: - errs = _count_col_errors(tbl, col["name"]) - col["error_count"] = errs - total_errors += errs - else: - col["error_count"] = 0 - - # Column-level dependency edges (within this table) - for col in columns: - if col["is_computed"] and col["computed_with"]: - deps = _parse_deps(col["computed_with"], all_col_names) - col["depends_on"] = deps - - # Indices - raw_indexes = md.get("indexes", {}) - indexes = [] - for idx_name, idx_info in raw_indexes.items(): - indexes.append( - { - "name": idx_name, - "columns": idx_info.get("columns", []), - "type": idx_info.get("index_type", "unknown"), - "embedding": str(idx_info.get("parameters", {}).get("embedding", ""))[:120], - } - ) + table_paths = sorted(pxt.list_tables(NAMESPACE, recursive=True)) - # Version history (last 10) - try: - raw_versions = tbl.get_versions() - versions = [] - for v in raw_versions[:10]: - versions.append( - { - "version": v["version"], - "created_at": v["created_at"].isoformat() if v.get("created_at") else None, - "change_type": v.get("change_type"), - "inserts": v.get("inserts", 0), - "updates": v.get("updates", 0), - "deletes": v.get("deletes", 0), - "errors": v.get("errors", 0), - } - ) - except Exception: - versions = [] + nodes: list[dict] = [] + edges: list[dict] = [] - base_path = md.get("base") - is_view = md.get("is_view", False) - - iterator_type = _detect_iterator(columns) if is_view else None - - nodes.append( + for path in table_paths: + try: + tbl = pxt.get_table(path) + md = tbl.get_metadata() + col_meta = md.get("columns", {}) + row_count = tbl.count() + + all_col_names = set(col_meta.keys()) + + columns = [] + computed_cols = [] + insertable_cols: set[str] = set() + + short_name = path.rsplit("/", 1)[-1] + + for col_name, info in col_meta.items(): + cw = info.get("computed_with") + is_computed = cw is not None + if not is_computed: + insertable_cols.add(col_name) + defined_in = info.get("defined_in") + + cw_str = str(cw)[:200] if cw else None + func_name = _extract_func_name(cw_str) if is_computed else None + func_type = _classify_func(func_name) if func_name else None + + col_entry: dict = { + "name": col_name, + "type": info.get("type_", "unknown"), + "is_computed": is_computed, + "computed_with": cw_str, + "defined_in": defined_in, + "defined_in_self": defined_in == short_name, + "func_name": func_name, + "func_type": func_type, + } + comment = info.get("comment") + if comment: + col_entry["comment"] = comment + custom_meta = info.get("custom_metadata") + if custom_meta: + col_entry["custom_metadata"] = custom_meta + columns.append(col_entry) + if is_computed: + computed_cols.append(col_name) + + # Compute error counts for computed columns (sample first 500 rows) + total_errors = 0 + for col in columns: + if col["is_computed"]: + errs = _count_col_errors(tbl, col["name"]) + col["error_count"] = errs + total_errors += errs + else: + col["error_count"] = 0 + + # Column-level dependency edges (within this table) + for col in columns: + if col["is_computed"] and col["computed_with"]: + deps = _parse_deps(col["computed_with"], all_col_names) + col["depends_on"] = deps + + # Indices + raw_indexes = md.get("indexes", {}) + indexes = [] + for idx_name, idx_info in raw_indexes.items(): + indexes.append( { - "path": path, - "name": short_name, - "is_view": is_view, - "base": base_path, - "row_count": row_count, - "version": md.get("version", 0), - "total_errors": total_errors, - "columns": columns, - "indexes": indexes, - "versions": versions, - "computed_count": len(computed_cols), - "insertable_count": len(columns) - len(computed_cols), - "iterator_type": iterator_type, + "name": idx_name, + "columns": idx_info.get("columns", []), + "type": idx_info.get("index_type", "unknown"), + "embedding": str(idx_info.get("parameters", {}).get("embedding", ""))[:120], } ) - if is_view and base_path: - edges.append( + # Version history (last 10) + try: + raw_versions = tbl.get_versions() + versions = [] + for v in raw_versions[:10]: + versions.append( { - "source": base_path, - "target": path, - "type": "view", - "label": iterator_type or "view", + "version": v["version"], + "created_at": v["created_at"].isoformat() if v.get("created_at") else None, + "change_type": v.get("change_type"), + "inserts": v.get("inserts", 0), + "updates": v.get("updates", 0), + "deletes": v.get("deletes", 0), + "errors": v.get("errors", 0), } ) + except Exception: + versions = [] - # Cross-table query edges (e.g., tools -> chunks via search_documents) - seen_query_targets: set[str] = set() - for col in columns: - fn = col.get("func_name") - if fn and fn in _QUERY_TABLE_MAP: - target_table = _QUERY_TABLE_MAP[fn] - edge_key = f"{path}->{target_table}" - if edge_key not in seen_query_targets: - seen_query_targets.add(edge_key) - edges.append( - { - "source": target_table, - "target": path, - "type": "query", - "label": fn, - } - ) - - except Exception as e: - logger.warning(f"Pipeline: could not inspect {path}: {e}") - nodes.append( - { - "path": path, - "name": path.split(".")[-1] if "." in path else path, - "is_view": False, - "base": None, - "row_count": 0, - "version": 0, - "total_errors": 0, - "columns": [], - "indexes": [], - "versions": [], - "computed_count": 0, - "insertable_count": 0, - "error": str(e), - } - ) - - return {"nodes": nodes, "edges": edges} - - except Exception as e: - logger.error(f"Pipeline error: {e}", exc_info=True) - raise HTTPException(status_code=500, detail=str(e)) + base_path = md.get("base") + is_view = md.get("is_view", False) + iterator_type = _detect_iterator(columns) if is_view else None -# ── Version History ─────────────────────────────────────────────────────────── + nodes.append( + { + "path": path, + "name": short_name, + "is_view": is_view, + "base": base_path, + "row_count": row_count, + "version": md.get("version", 0), + "total_errors": total_errors, + "columns": columns, + "indexes": indexes, + "versions": versions, + "computed_count": len(computed_cols), + "insertable_count": len(columns) - len(computed_cols), + "iterator_type": iterator_type, + } + ) + if is_view and base_path: + edges.append( + { + "source": base_path, + "target": path, + "type": "view", + "label": iterator_type or "view", + } + ) -@router.get("/table/{path:path}/versions") -@pxt_retry() -def get_table_versions(path: str, limit: int = 20): - """Get the version history for a specific table.""" - path = require_allowed_table(path) - try: - tbl = pxt.get_table(path) - except Exception: - raise HTTPException(status_code=404, detail=f"Table '{path}' not found") + # Cross-table query edges (e.g., tools -> chunks via search_documents) + seen_query_targets: set[str] = set() + for col in columns: + fn = col.get("func_name") + if fn and fn in _QUERY_TABLE_MAP: + target_table = _QUERY_TABLE_MAP[fn] + edge_key = f"{path}->{target_table}" + if edge_key not in seen_query_targets: + seen_query_targets.add(edge_key) + edges.append( + { + "source": target_table, + "target": path, + "type": "query", + "label": fn, + } + ) - try: - raw_versions = tbl.get_versions() - versions = [] - for v in raw_versions[:limit]: - versions.append( + except Exception as e: + logger.warning(f"Pipeline: could not inspect {path}: {e}") + nodes.append( { - "version": v["version"], - "created_at": v["created_at"].isoformat() if v.get("created_at") else None, - "change_type": v.get("change_type"), - "inserts": v.get("inserts", 0), - "updates": v.get("updates", 0), - "deletes": v.get("deletes", 0), - "errors": v.get("errors", 0), - "schema_change": v.get("schema_change"), + "path": path, + "name": path.split(".")[-1] if "." in path else path, + "is_view": False, + "base": None, + "row_count": 0, + "version": 0, + "total_errors": 0, + "columns": [], + "indexes": [], + "versions": [], + "computed_count": 0, + "insertable_count": 0, + "error": str(e), } ) - return { - "path": path, - "current_version": versions[0]["version"] if versions else 0, - "can_revert": len(versions) > 1, - "versions": versions, - } - except Exception as e: - logger.error(f"get_versions error for {path}: {e}", exc_info=True) - raise HTTPException(status_code=500, detail=str(e)) + return {"nodes": nodes, "edges": edges} diff --git a/backend/pixelbot/routers/experiments.py b/backend/pixelbot/routers/experiments.py index 68b82a5..64a4873 100644 --- a/backend/pixelbot/routers/experiments.py +++ b/backend/pixelbot/routers/experiments.py @@ -249,7 +249,6 @@ def get_available_models(): @router.post("/run", response_model=RunExperimentResponse) -@pxt_retry() def run_experiment(body: RunExperimentRequest): """Run a prompt against multiple models in parallel and store results.""" if not body.user_prompt.strip(): @@ -419,75 +418,64 @@ def get_experiment_history(): @pxt_retry() def get_experiment(experiment_id: str): """Return full results for a specific experiment.""" - try: - table = pxt.get_table(_TABLE_PATH) - rows = ( - table.where((table.experiment_id == experiment_id) & (table.user_id == config.DEFAULT_USER_ID)) - .select( - table.experiment_id, - table.task, - table.system_prompt, - table.user_prompt, - table.model_id, - table.model_name, - table.provider, - table.temperature, - table.max_tokens, - table.response, - table.response_time_ms, - table.word_count, - table.char_count, - table.error, - table.timestamp, - ) - .collect() + table = pxt.get_table(_TABLE_PATH) + rows = ( + table.where((table.experiment_id == experiment_id) & (table.user_id == config.DEFAULT_USER_ID)) + .select( + table.experiment_id, + table.task, + table.system_prompt, + table.user_prompt, + table.model_id, + table.model_name, + table.provider, + table.temperature, + table.max_tokens, + table.response, + table.response_time_ms, + table.word_count, + table.char_count, + table.error, + table.timestamp, ) + .collect() + ) - if not rows: - raise HTTPException(status_code=404, detail="Experiment not found") - - first = rows[0] - results = [] - for row in rows: - error_val = row.get("error") or "" - results.append( - ExperimentResult( - model_id=row["model_id"], - model_name=row["model_name"], - provider=row["provider"], - response=row["response"] if not error_val else None, - response_time_ms=row.get("response_time_ms", 0), - word_count=row.get("word_count", 0), - char_count=row.get("char_count", 0), - error=error_val if error_val else None, - ) + if not rows: + raise HTTPException(status_code=404, detail="Experiment not found") + + first = rows[0] + results = [] + for row in rows: + error_val = row.get("error") or "" + results.append( + ExperimentResult( + model_id=row["model_id"], + model_name=row["model_name"], + provider=row["provider"], + response=row["response"] if not error_val else None, + response_time_ms=row.get("response_time_ms", 0), + word_count=row.get("word_count", 0), + char_count=row.get("char_count", 0), + error=error_val if error_val else None, ) - - return RunExperimentResponse( - experiment_id=experiment_id, - task=first["task"], - system_prompt=first["system_prompt"], - user_prompt=first["user_prompt"], - temperature=first["temperature"], - max_tokens=first["max_tokens"], - timestamp=first["timestamp"].isoformat() if first["timestamp"] else "", - results=results, ) - except HTTPException: - raise - except Exception as e: - logger.error(f"Failed to get experiment {experiment_id}: {e}", exc_info=True) - raise HTTPException(status_code=500, detail=str(e)) + + return RunExperimentResponse( + experiment_id=experiment_id, + task=first["task"], + system_prompt=first["system_prompt"], + user_prompt=first["user_prompt"], + temperature=first["temperature"], + max_tokens=first["max_tokens"], + timestamp=first["timestamp"].isoformat() if first["timestamp"] else "", + results=results, + ) @router.delete("/{experiment_id}") -@pxt_retry() def delete_experiment(experiment_id: str): """Delete all results for an experiment.""" - try: - table = pxt.get_table(_TABLE_PATH) - table.delete(where=(table.experiment_id == experiment_id) & (table.user_id == config.DEFAULT_USER_ID)) - return {"message": f"Experiment {experiment_id} deleted"} - except Exception as e: - logger.error(f"Failed to delete experiment {experiment_id}: {e}", exc_info=True) - raise HTTPException(status_code=500, detail=str(e)) + table = pxt.get_table(_TABLE_PATH) + table.delete(where=(table.experiment_id == experiment_id) & (table.user_id == config.DEFAULT_USER_ID)) + return {"message": f"Experiment {experiment_id} deleted"} diff --git a/backend/pixelbot/routers/export.py b/backend/pixelbot/routers/export.py index 4572358..672af92 100644 --- a/backend/pixelbot/routers/export.py +++ b/backend/pixelbot/routers/export.py @@ -68,28 +68,24 @@ def _collect_rows(table_path: str, limit: int, columns: list[str] | None) -> tup @pxt_retry() def list_exportable_tables(): """Return all tables with their column info for the export picker.""" - try: - tables_raw = list(pxt.list_tables(config.APP_NAMESPACE, recursive=True)) - tables_raw.extend(registered_scratch_tables()) - result = [] - for path in tables_raw: - try: - tbl = pxt.get_table(path) - col_names = tbl.columns() - row_count = tbl.count() - result.append( - { - "path": path, - "columns": col_names, - "row_count": row_count, - } - ) - except Exception: - result.append({"path": path, "columns": [], "row_count": 0}) - return {"tables": result} - except Exception as e: - logger.error(f"Failed to list tables: {e}", exc_info=True) - raise HTTPException(status_code=500, detail=str(e)) + tables_raw = list(pxt.list_tables(config.APP_NAMESPACE, recursive=True)) + tables_raw.extend(registered_scratch_tables()) + result = [] + for path in tables_raw: + try: + tbl = pxt.get_table(path) + col_names = tbl.columns() + row_count = tbl.count() + result.append( + { + "path": path, + "columns": col_names, + "row_count": row_count, + } + ) + except Exception: + result.append({"path": path, "columns": [], "row_count": 0}) + return {"tables": result} # ── Export as JSON ─────────────────────────────────────────────────────────── @@ -230,9 +226,6 @@ def export_json_column( } except ImportError: raise HTTPException(status_code=501, detail="pixeltable.functions.json not available") - except Exception as e: - logger.error(f"JSON dumps error: {e}", exc_info=True) - raise HTTPException(status_code=500, detail=str(e)) @router.get("/preview/{table_path:path}", response_model=PreviewResponse) diff --git a/backend/pixelbot/routers/files.py b/backend/pixelbot/routers/files.py index f05128f..8f87d4e 100644 --- a/backend/pixelbot/routers/files.py +++ b/backend/pixelbot/routers/files.py @@ -18,8 +18,6 @@ MEDIA_ROW_MODELS, AddUrlResponse, CsvRegistryRow, - DeleteAllResponse, - DeleteFileResponse, UploadResponse, ) from pixelbot.utils import create_thumbnail_base64, pxt_retry @@ -166,7 +164,6 @@ def _source_to_filename(source) -> str: @router.post("/upload", response_model=UploadResponse) -@pxt_retry() def upload_file(file: UploadFile = File(...)): """Handle file uploads. CSVs are imported into their own Pixeltable table.""" user_id = config.DEFAULT_USER_ID @@ -200,9 +197,6 @@ def upload_file(file: UploadFile = File(...)): except FileNotFoundError: pass raise - except Exception as e: - logger.error(f"Error saving file to disk: {e}", exc_info=True) - raise HTTPException(status_code=500, detail=str(e)) # CSV files get their own Pixeltable table if file_ext == "csv": @@ -213,27 +207,21 @@ def upload_file(file: UploadFile = File(...)): raise HTTPException(status_code=400, detail=f"Unsupported file extension: {file_ext}") table_key, data_col = mapping + file_uuid = str(uuid.uuid4()) + current_timestamp = datetime.now() - try: - file_uuid = str(uuid.uuid4()) - current_timestamp = datetime.now() + table = get_pxt_table(table_key) + RowModel = MEDIA_ROW_MODELS[table_key] + row = RowModel(**{data_col: file_path, "uuid": file_uuid, "timestamp": current_timestamp, "user_id": user_id}) + status = table.insert([row], return_rows=True) + if status.errors: + raise RuntimeError(f"Insert failed: {status.errors}") - table = get_pxt_table(table_key) - RowModel = MEDIA_ROW_MODELS[table_key] - row = RowModel(**{data_col: file_path, "uuid": file_uuid, "timestamp": current_timestamp, "user_id": user_id}) - status = table.insert([row], return_rows=True) - if status.errors: - raise RuntimeError(f"Insert failed: {status.errors}") - - return UploadResponse( - message=f"File successfully uploaded to {table_key} table", - filename=safe_name, - uuid=file_uuid, - ) - - except Exception as e: - logger.error(f"Error uploading file: {e}", exc_info=True) - raise HTTPException(status_code=500, detail=str(e)) + return UploadResponse( + message=f"File successfully uploaded to {table_key} table", + filename=safe_name, + uuid=file_uuid, + ) def _import_csv(file_path: str, display_name: str, user_id: str) -> UploadResponse: @@ -285,7 +273,7 @@ def _import_csv(file_path: str, display_name: str, user_id: str) -> UploadRespon except Exception: pass logger.error(f"Error importing CSV: {e}", exc_info=True) - raise HTTPException(status_code=500, detail=str(e)) + raise # ── Add URL ─────────────────────────────────────────────────────────────────── @@ -296,7 +284,6 @@ class AddUrlRequest(BaseModel): @router.post("/add_url", response_model=AddUrlResponse) -@pxt_retry() def add_url(body: AddUrlRequest): """Add a URL as a data source.""" user_id = config.DEFAULT_USER_ID @@ -343,100 +330,7 @@ def add_url(body: AddUrlRequest): detail="Document is too large to process (exceeds 1M characters). Try a shorter document or a direct file upload.", ) logger.error(f"Error adding URL: {e}", exc_info=True) - raise HTTPException(status_code=500, detail=str(e)) - except Exception as e: - logger.error(f"Error adding URL: {e}", exc_info=True) - raise HTTPException(status_code=500, detail=str(e)) - - -# ── Delete File ─────────────────────────────────────────────────────────────── - - -@router.delete("/delete_file/{file_uuid}/{file_type}", response_model=DeleteFileResponse) -@pxt_retry() -def delete_file(file_uuid: str, file_type: str): - """Delete a file by UUID and type.""" - user_id = config.DEFAULT_USER_ID - - if file_type not in TABLE_MAP: - raise HTTPException(status_code=400, detail=f"Invalid file type: {file_type}") - - try: - table = get_pxt_table(file_type) - data_col_map = {"document": "document", "image": "image", "video": "video", "audio": "audio"} - data_col = data_col_map.get(file_type) - if not data_col: - raise HTTPException(status_code=400, detail=f"Cannot map file_type '{file_type}'") - - # Retrieve file path before deletion - file_path_to_delete = None - try: - record = ( - table.where((table.uuid == file_uuid) & (table.user_id == user_id)) - .select(file_source=getattr(table, data_col)) - .collect() - ) - if len(record) > 0: - file_source = record[0].get("file_source") - if isinstance(file_source, str) and not file_source.startswith(("http://", "https://")): - possible = os.path.abspath(os.path.join(config.UPLOAD_FOLDER, os.path.basename(file_source))) - if os.path.exists(possible): - file_path_to_delete = possible - except Exception as e: - logger.error(f"Error retrieving file path: {e}") - - # Delete from DB - status = table.delete(where=(table.uuid == file_uuid) & (table.user_id == user_id)) - db_deleted = status.num_rows > 0 - - file_deleted = False - if db_deleted and file_path_to_delete: - try: - os.remove(file_path_to_delete) - file_deleted = True - except Exception as e: - logger.error(f"Error deleting file from disk: {e}") - - if not db_deleted: - raise HTTPException(status_code=404, detail=f"No {file_type} found with UUID {file_uuid}") - - return DeleteFileResponse( - message=f"{file_type.capitalize()} deleted successfully", - db_deleted=db_deleted, - file_deleted=file_deleted, - uuid=file_uuid, - ) - - except HTTPException: raise - except Exception as e: - logger.error(f"Error deleting file: {e}", exc_info=True) - raise HTTPException(status_code=500, detail=str(e)) - - -# ── Delete All ──────────────────────────────────────────────────────────────── - - -class DeleteAllRequest(BaseModel): - type: str - - -@router.post("/delete_all", response_model=DeleteAllResponse) -@pxt_retry() -def delete_all(body: DeleteAllRequest): - """Delete all items from a given table type.""" - user_id = config.DEFAULT_USER_ID - - if body.type not in TABLE_MAP: - raise HTTPException(status_code=400, detail=f"Invalid type. Must be one of: {', '.join(TABLE_MAP.keys())}") - - try: - table = get_pxt_table(body.type) - status = table.delete(where=table.user_id == user_id) - return DeleteAllResponse(message=f"Deleted {status.num_rows} {body.type} items") - except Exception as e: - logger.error(f"Error deleting all {body.type}: {e}", exc_info=True) - raise HTTPException(status_code=500, detail=str(e)) # ── Context Info ────────────────────────────────────────────────────────────── @@ -462,175 +356,169 @@ def get_context_info(): direct ResultSet iteration (no pandas conversion). """ user_id = config.DEFAULT_USER_ID + # Available tools + available_tools = [ + {"name": "get_latest_news", "description": inspect.getdoc(functions.get_latest_news)}, + {"name": "fetch_financial_data", "description": inspect.getdoc(functions.fetch_financial_data)}, + {"name": "search_news", "description": inspect.getdoc(functions.search_news)}, + ] + + # Documents — direct iteration over ResultSet + document_list: list[dict] = [] + try: + doc_table = get_pxt_table("document") + for row in ( + doc_table.where(doc_table.user_id == user_id) + .select(doc_source=doc_table.document, uuid_col=doc_table.uuid) + .collect() + ): + document_list.append({"name": _source_to_filename(row["doc_source"]), "uuid": row["uuid_col"]}) + except Exception as e: + logger.error(f"Error fetching documents: {e}") + # Images — use precomputed `thumbnail` column from Pixeltable + image_list: list[dict] = [] try: - # Available tools - available_tools = [ - {"name": "get_latest_news", "description": inspect.getdoc(functions.get_latest_news)}, - {"name": "fetch_financial_data", "description": inspect.getdoc(functions.fetch_financial_data)}, - {"name": "search_news", "description": inspect.getdoc(functions.search_news)}, - ] - - # Documents — direct iteration over ResultSet - document_list: list[dict] = [] - try: - doc_table = get_pxt_table("document") - for row in ( - doc_table.where(doc_table.user_id == user_id) - .select(doc_source=doc_table.document, uuid_col=doc_table.uuid) - .collect() - ): - document_list.append({"name": _source_to_filename(row["doc_source"]), "uuid": row["uuid_col"]}) - except Exception as e: - logger.error(f"Error fetching documents: {e}") + img_table = get_pxt_table("image") + for row in ( + img_table.where(img_table.user_id == user_id) + .select( + img_source=img_table.image, + uuid_col=img_table.uuid, + thumb=img_table.thumbnail, + ) + .collect() + ): + thumbnail = _pxt_thumbnail_to_data_uri(row.get("thumb")) + image_list.append( + { + "name": _source_to_filename(row["img_source"]), + "thumbnail": thumbnail, + "uuid": row["uuid_col"], + } + ) + except Exception as e: + logger.error(f"Error fetching images: {e}") - # Images — use precomputed `thumbnail` column from Pixeltable - image_list: list[dict] = [] - try: - img_table = get_pxt_table("image") - for row in ( - img_table.where(img_table.user_id == user_id) - .select( - img_source=img_table.image, - uuid_col=img_table.uuid, - thumb=img_table.thumbnail, - ) - .collect() - ): - thumbnail = _pxt_thumbnail_to_data_uri(row.get("thumb")) - image_list.append( - { - "name": _source_to_filename(row["img_source"]), - "thumbnail": thumbnail, - "uuid": row["uuid_col"], - } - ) - except Exception as e: - logger.error(f"Error fetching images: {e}") + # Videos (with thumbnails from first frame) + video_list: list[dict] = [] + try: + vid_table = get_pxt_table("video") + video_frames_view = pxt.get_table("pixelbot_v3.video_frames") - # Videos (with thumbnails from first frame) - video_list: list[dict] = [] + # Build a map of uuid → first-frame thumbnail + first_frames_map: dict[str, str | None] = {} try: - vid_table = get_pxt_table("video") - video_frames_view = pxt.get_table("pixelbot_v3.video_frames") - - # Build a map of uuid → first-frame thumbnail - first_frames_map: dict[str, str | None] = {} - try: - for row in ( - video_frames_view.where(video_frames_view.pos == 0) - .select( - video_uuid=video_frames_view.uuid, - frame=video_frames_view.frame, - ) - .collect() - ): - frame = row.get("frame") - if isinstance(frame, Image.Image): - first_frames_map[row["video_uuid"]] = create_thumbnail_base64(frame, THUMB_SIZE_SIDEBAR) - except Exception as e: - logger.error(f"Error fetching video first frames: {e}") - for row in ( - vid_table.where(vid_table.user_id == user_id) + video_frames_view.where(video_frames_view.pos == 0) .select( - video_col=vid_table.video, - uuid_col=vid_table.uuid, + video_uuid=video_frames_view.uuid, + frame=video_frames_view.frame, ) .collect() ): - video_list.append( - { - "name": _source_to_filename(row["video_col"]), - "thumbnail": first_frames_map.get(row["uuid_col"]), - "uuid": row["uuid_col"], - } - ) + frame = row.get("frame") + if isinstance(frame, Image.Image): + first_frames_map[row["video_uuid"]] = create_thumbnail_base64(frame, THUMB_SIZE_SIDEBAR) except Exception as e: - logger.error(f"Error fetching videos: {e}") + logger.error(f"Error fetching video first frames: {e}") - # Audios - audio_list: list[dict] = [] - try: - audio_table = get_pxt_table("audio") - for row in ( - audio_table.where(audio_table.user_id == user_id) - .select( - audio_col=audio_table.audio, - uuid_col=audio_table.uuid, - ) - .collect() - ): - audio_list.append({"name": _source_to_filename(row["audio_col"]), "uuid": row["uuid_col"]}) - except Exception as e: - logger.error(f"Error fetching audios: {e}") + for row in ( + vid_table.where(vid_table.user_id == user_id) + .select( + video_col=vid_table.video, + uuid_col=vid_table.uuid, + ) + .collect() + ): + video_list.append( + { + "name": _source_to_filename(row["video_col"]), + "thumbnail": first_frames_map.get(row["uuid_col"]), + "uuid": row["uuid_col"], + } + ) + except Exception as e: + logger.error(f"Error fetching videos: {e}") - # CSV tables (from registry) - csv_tables: list[dict] = [] - try: - csv_registry = pxt.get_table("pixelbot_v3.csv_registry") - for row in ( - csv_registry.where(csv_registry.user_id == user_id) - .select( - csv_registry.display_name, - csv_registry.uuid, - csv_registry.row_count, - csv_registry.col_names, - ) - .collect() - ): - csv_tables.append( - { - "name": row["display_name"], - "uuid": row["uuid"], - "row_count": row["row_count"], - "columns": row["col_names"], - } - ) - except Exception as e: - logger.error(f"Error fetching CSV tables: {e}") + # Audios + audio_list: list[dict] = [] + try: + audio_table = get_pxt_table("audio") + for row in ( + audio_table.where(audio_table.user_id == user_id) + .select( + audio_col=audio_table.audio, + uuid_col=audio_table.uuid, + ) + .collect() + ): + audio_list.append({"name": _source_to_filename(row["audio_col"]), "uuid": row["uuid_col"]}) + except Exception as e: + logger.error(f"Error fetching audios: {e}") - # Workflow history — direct iteration, no pandas - workflow_data: list[dict] = [] - try: - wf_table = pxt.get_table("pixelbot_v3.tools") - for row in ( - wf_table.where(wf_table.user_id == user_id) - .select( - wf_table.timestamp, - wf_table.prompt, - wf_table.answer, - ) - .order_by(wf_table.timestamp, asc=False) - .collect() - ): - ts = row.get("timestamp") - workflow_data.append( - { - "timestamp": ts.strftime("%Y-%m-%d %H:%M:%S.%f") if ts else None, - "prompt": row.get("prompt"), - "answer": row.get("answer"), - } - ) - except Exception as e: - logger.error(f"Error fetching workflow data: {e}") - - return { - "tools": available_tools, - "documents": document_list, - "images": image_list, - "videos": video_list, - "audios": audio_list, - "csv_tables": csv_tables, - "initial_prompt": config.INITIAL_SYSTEM_PROMPT, - "final_prompt": config.FINAL_SYSTEM_PROMPT, - "workflow_data": workflow_data, - "parameters": { - "max_tokens": config.DEFAULT_MAX_TOKENS, - "temperature": config.DEFAULT_TEMPERATURE, - }, - } + # CSV tables (from registry) + csv_tables: list[dict] = [] + try: + csv_registry = pxt.get_table("pixelbot_v3.csv_registry") + for row in ( + csv_registry.where(csv_registry.user_id == user_id) + .select( + csv_registry.display_name, + csv_registry.uuid, + csv_registry.row_count, + csv_registry.col_names, + ) + .collect() + ): + csv_tables.append( + { + "name": row["display_name"], + "uuid": row["uuid"], + "row_count": row["row_count"], + "columns": row["col_names"], + } + ) + except Exception as e: + logger.error(f"Error fetching CSV tables: {e}") + # Workflow history — direct iteration, no pandas + workflow_data: list[dict] = [] + try: + wf_table = pxt.get_table("pixelbot_v3.tools") + for row in ( + wf_table.where(wf_table.user_id == user_id) + .select( + wf_table.timestamp, + wf_table.prompt, + wf_table.answer, + ) + .order_by(wf_table.timestamp, asc=False) + .collect() + ): + ts = row.get("timestamp") + workflow_data.append( + { + "timestamp": ts.strftime("%Y-%m-%d %H:%M:%S.%f") if ts else None, + "prompt": row.get("prompt"), + "answer": row.get("answer"), + } + ) except Exception as e: - logger.error(f"Error fetching context info: {e}", exc_info=True) - raise HTTPException(status_code=500, detail=str(e)) + logger.error(f"Error fetching workflow data: {e}") + + return { + "tools": available_tools, + "documents": document_list, + "images": image_list, + "videos": video_list, + "audios": audio_list, + "csv_tables": csv_tables, + "initial_prompt": config.INITIAL_SYSTEM_PROMPT, + "final_prompt": config.FINAL_SYSTEM_PROMPT, + "workflow_data": workflow_data, + "parameters": { + "max_tokens": config.DEFAULT_MAX_TOKENS, + "temperature": config.DEFAULT_TEMPERATURE, + }, + } diff --git a/backend/pixelbot/routers/history.py b/backend/pixelbot/routers/history.py index 0a11719..4fb4bb4 100644 --- a/backend/pixelbot/routers/history.py +++ b/backend/pixelbot/routers/history.py @@ -4,7 +4,7 @@ from datetime import datetime import pixeltable as pxt -from fastapi import APIRouter, HTTPException +from fastapi import APIRouter from fastapi.responses import StreamingResponse from pixelbot import config @@ -23,49 +23,43 @@ def list_conversations(): """List all conversations, grouped by conversation_id.""" user_id = config.DEFAULT_USER_ID - - try: - table = pxt.get_table("pixelbot_v3.chat_history") - rows = list( - table.where(table.user_id == user_id) - .select( - role=table.role, - content=table.content, - conversation_id=table.conversation_id, - timestamp=table.timestamp, - ) - .order_by(table.timestamp, asc=True) - .collect() + table = pxt.get_table("pixelbot_v3.chat_history") + rows = list( + table.where(table.user_id == user_id) + .select( + role=table.role, + content=table.content, + conversation_id=table.conversation_id, + timestamp=table.timestamp, ) - - convos: dict[str, dict] = {} - for row in rows: - cid = row.get("conversation_id") or "default" - if cid not in convos: - convos[cid] = { - "conversation_id": cid, - "title": "", - "created_at": row["timestamp"].isoformat() - if isinstance(row["timestamp"], datetime) - else str(row["timestamp"]), - "updated_at": row["timestamp"].isoformat() - if isinstance(row["timestamp"], datetime) - else str(row["timestamp"]), - "message_count": 0, - } - entry = convos[cid] - entry["message_count"] += 1 - ts_iso = row["timestamp"].isoformat() if isinstance(row["timestamp"], datetime) else str(row["timestamp"]) - entry["updated_at"] = ts_iso - if not entry["title"] and row["role"] == "user": - entry["title"] = row["content"][:100] - - result = sorted(convos.values(), key=lambda c: c["updated_at"], reverse=True) - return result - - except Exception as e: - logger.error(f"Error listing conversations: {e}", exc_info=True) - raise HTTPException(status_code=500, detail=str(e)) + .order_by(table.timestamp, asc=True) + .collect() + ) + + convos: dict[str, dict] = {} + for row in rows: + cid = row.get("conversation_id") or "default" + if cid not in convos: + convos[cid] = { + "conversation_id": cid, + "title": "", + "created_at": row["timestamp"].isoformat() + if isinstance(row["timestamp"], datetime) + else str(row["timestamp"]), + "updated_at": row["timestamp"].isoformat() + if isinstance(row["timestamp"], datetime) + else str(row["timestamp"]), + "message_count": 0, + } + entry = convos[cid] + entry["message_count"] += 1 + ts_iso = row["timestamp"].isoformat() if isinstance(row["timestamp"], datetime) else str(row["timestamp"]) + entry["updated_at"] = ts_iso + if not entry["title"] and row["role"] == "user": + entry["title"] = row["content"][:100] + + result = sorted(convos.values(), key=lambda c: c["updated_at"], reverse=True) + return result @router.get("/conversations/{conversation_id}", response_model=ConversationDetail) @@ -73,130 +67,36 @@ def list_conversations(): def get_conversation(conversation_id: str): """Get all messages for a specific conversation.""" user_id = config.DEFAULT_USER_ID - - try: - table = pxt.get_table("pixelbot_v3.chat_history") - rows = list( - table.where((table.user_id == user_id) & (table.conversation_id == conversation_id)) - .select(role=table.role, content=table.content, timestamp=table.timestamp) - .order_by(table.timestamp, asc=True) - .collect() + table = pxt.get_table("pixelbot_v3.chat_history") + rows = list( + table.where((table.user_id == user_id) & (table.conversation_id == conversation_id)) + .select(role=table.role, content=table.content, timestamp=table.timestamp) + .order_by(table.timestamp, asc=True) + .collect() + ) + + messages = [] + for row in rows: + messages.append( + { + "role": row["role"], + "content": row["content"], + "timestamp": row["timestamp"].isoformat() + if isinstance(row["timestamp"], datetime) + else str(row["timestamp"]), + } ) - messages = [] - for row in rows: - messages.append( - { - "role": row["role"], - "content": row["content"], - "timestamp": row["timestamp"].isoformat() - if isinstance(row["timestamp"], datetime) - else str(row["timestamp"]), - } - ) - - return {"conversation_id": conversation_id, "messages": messages} - - except Exception as e: - logger.error(f"Error fetching conversation: {e}", exc_info=True) - raise HTTPException(status_code=500, detail=str(e)) + return {"conversation_id": conversation_id, "messages": messages} @router.delete("/conversations/{conversation_id}", response_model=DeleteResponse) -@pxt_retry() def delete_conversation(conversation_id: str): """Delete all messages in a conversation.""" user_id = config.DEFAULT_USER_ID - - try: - table = pxt.get_table("pixelbot_v3.chat_history") - status = table.delete(where=(table.user_id == user_id) & (table.conversation_id == conversation_id)) - return {"message": "Conversation deleted", "num_deleted": status.num_rows} - - except Exception as e: - logger.error(f"Error deleting conversation: {e}", exc_info=True) - raise HTTPException(status_code=500, detail=str(e)) - - -def _parse_timestamp(ts_str: str) -> datetime: - """Parse a timestamp string, trying multiple formats.""" - for fmt in ("%Y-%m-%d %H:%M:%S.%f", "%Y-%m-%d %H:%M:%S"): - try: - return datetime.strptime(ts_str, fmt) - except ValueError: - continue - raise HTTPException(status_code=400, detail="Invalid timestamp format. Expected YYYY-MM-DD HH:MM:SS[.ffffff]") - - -# ── Workflow Detail ─────────────────────────────────────────────────────────── - - -@router.get("/workflow_detail/{timestamp_str:path}") -@pxt_retry() -def get_workflow_detail(timestamp_str: str): - """Get full detail for a specific workflow entry.""" - user_id = config.DEFAULT_USER_ID - target_timestamp = _parse_timestamp(timestamp_str) - - try: - wf_table = pxt.get_table("pixelbot_v3.tools") - result_df = ( - wf_table.where((wf_table.timestamp == target_timestamp) & (wf_table.user_id == user_id)) - .select( - prompt=wf_table.prompt, - timestamp=wf_table.timestamp, - initial_system_prompt=wf_table.initial_system_prompt, - final_system_prompt=wf_table.final_system_prompt, - initial_response=wf_table.initial_response, - tool_output=wf_table.tool_output, - final_response=wf_table.final_response, - answer=wf_table.answer, - max_tokens=wf_table.max_tokens, - temperature=wf_table.temperature, - ) - .collect() - ) - - if len(result_df) == 0: - raise HTTPException(status_code=404, detail="Workflow entry not found") - - detail = result_df[0] - if "timestamp" in detail and isinstance(detail["timestamp"], datetime): - detail["timestamp"] = detail["timestamp"].isoformat() - - return detail - - except HTTPException: - raise - except Exception as e: - logger.error(f"Error fetching workflow detail: {e}", exc_info=True) - raise HTTPException(status_code=500, detail=str(e)) - - -# ── Delete History Entry ────────────────────────────────────────────────────── - - -@router.delete("/delete_history/{timestamp_str:path}", response_model=DeleteResponse) -@pxt_retry() -def delete_history_entry(timestamp_str: str): - """Delete a specific history entry by timestamp.""" - user_id = config.DEFAULT_USER_ID - target_timestamp = _parse_timestamp(timestamp_str) - - try: - wf_table = pxt.get_table("pixelbot_v3.tools") - status = wf_table.delete(where=(wf_table.timestamp == target_timestamp) & (wf_table.user_id == user_id)) - - if status.num_rows == 0: - raise HTTPException(status_code=404, detail="No entry found with that timestamp") - - return DeleteResponse(message="History entry deleted", num_deleted=status.num_rows) - - except HTTPException: - raise - except Exception as e: - logger.error(f"Error deleting history entry: {e}", exc_info=True) - raise HTTPException(status_code=500, detail=str(e)) + table = pxt.get_table("pixelbot_v3.chat_history") + status = table.delete(where=(table.user_id == user_id) & (table.conversation_id == conversation_id)) + return {"message": "Conversation deleted", "num_deleted": status.num_rows} # ── Download History ────────────────────────────────────────────────────────── @@ -207,40 +107,34 @@ def delete_history_entry(timestamp_str: str): def download_chat_history(): """Download the full chat history as JSON (direct iteration, no pandas).""" user_id = config.DEFAULT_USER_ID - - try: - wf_table = pxt.get_table("pixelbot_v3.tools") - rows = list( - wf_table.where(wf_table.user_id == user_id) - .select( - prompt=wf_table.prompt, - timestamp=wf_table.timestamp, - answer=wf_table.answer, - initial_system_prompt=wf_table.initial_system_prompt, - final_system_prompt=wf_table.final_system_prompt, - max_tokens=wf_table.max_tokens, - temperature=wf_table.temperature, - ) - .order_by(wf_table.timestamp, asc=False) - .collect() + wf_table = pxt.get_table("pixelbot_v3.tools") + rows = list( + wf_table.where(wf_table.user_id == user_id) + .select( + prompt=wf_table.prompt, + timestamp=wf_table.timestamp, + answer=wf_table.answer, + initial_system_prompt=wf_table.initial_system_prompt, + final_system_prompt=wf_table.final_system_prompt, + max_tokens=wf_table.max_tokens, + temperature=wf_table.temperature, ) + .order_by(wf_table.timestamp, asc=False) + .collect() + ) - for row in rows: - ts = row.get("timestamp") - if ts: - row["timestamp"] = ts.isoformat() + for row in rows: + ts = row.get("timestamp") + if ts: + row["timestamp"] = ts.isoformat() - json_bytes = json.dumps(rows, indent=2, default=str).encode("utf-8") + json_bytes = json.dumps(rows, indent=2, default=str).encode("utf-8") - return StreamingResponse( - io.BytesIO(json_bytes), - media_type="application/json", - headers={"Content-Disposition": "attachment; filename=chat_history_full.json"}, - ) - - except Exception as e: - logger.error(f"Error downloading history: {e}", exc_info=True) - raise HTTPException(status_code=500, detail=str(e)) + return StreamingResponse( + io.BytesIO(json_bytes), + media_type="application/json", + headers={"Content-Disposition": "attachment; filename=chat_history_full.json"}, + ) # ── Debug Export (full pixelbot_v3.tools table) ──────────────────────────────────── @@ -266,48 +160,42 @@ def _safe_serialize(obj: object) -> object: def debug_export(): """Export the full pixelbot_v3.tools table with every column for debugging.""" user_id = config.DEFAULT_USER_ID - - try: - wf_table = pxt.get_table("pixelbot_v3.tools") - - rows = list( - wf_table.where(wf_table.user_id == user_id) - .select( - prompt=wf_table.prompt, - timestamp=wf_table.timestamp, - user_id=wf_table.user_id, - initial_system_prompt=wf_table.initial_system_prompt, - final_system_prompt=wf_table.final_system_prompt, - max_tokens=wf_table.max_tokens, - temperature=wf_table.temperature, - initial_response=wf_table.initial_response, - tool_output=wf_table.tool_output, - doc_context=wf_table.doc_context, - image_context=wf_table.image_context, - video_frame_context=wf_table.video_frame_context, - memory_context=wf_table.memory_context, - chat_memory_context=wf_table.chat_memory_context, - history_context=wf_table.history_context, - multimodal_context_summary=wf_table.multimodal_context_summary, - final_prompt_messages=wf_table.final_prompt_messages, - final_response=wf_table.final_response, - answer=wf_table.answer, - follow_up_input_message=wf_table.follow_up_input_message, - follow_up_text=wf_table.follow_up_text, - ) - .order_by(wf_table.timestamp, asc=False) - .collect() + wf_table = pxt.get_table("pixelbot_v3.tools") + + rows = list( + wf_table.where(wf_table.user_id == user_id) + .select( + prompt=wf_table.prompt, + timestamp=wf_table.timestamp, + user_id=wf_table.user_id, + initial_system_prompt=wf_table.initial_system_prompt, + final_system_prompt=wf_table.final_system_prompt, + max_tokens=wf_table.max_tokens, + temperature=wf_table.temperature, + initial_response=wf_table.initial_response, + tool_output=wf_table.tool_output, + doc_context=wf_table.doc_context, + image_context=wf_table.image_context, + video_frame_context=wf_table.video_frame_context, + memory_context=wf_table.memory_context, + chat_memory_context=wf_table.chat_memory_context, + history_context=wf_table.history_context, + multimodal_context_summary=wf_table.multimodal_context_summary, + final_prompt_messages=wf_table.final_prompt_messages, + final_response=wf_table.final_response, + answer=wf_table.answer, + follow_up_input_message=wf_table.follow_up_input_message, + follow_up_text=wf_table.follow_up_text, ) - - sanitized = [_safe_serialize(row) for row in rows] - json_bytes = json.dumps(sanitized, indent=2, default=str).encode("utf-8") - - return StreamingResponse( - io.BytesIO(json_bytes), - media_type="application/json", - headers={"Content-Disposition": "attachment; filename=agents_tools_debug_export.json"}, - ) - - except Exception as e: - logger.error(f"Error during debug export: {e}", exc_info=True) - raise HTTPException(status_code=500, detail=str(e)) + .order_by(wf_table.timestamp, asc=False) + .collect() + ) + + sanitized = [_safe_serialize(row) for row in rows] + json_bytes = json.dumps(sanitized, indent=2, default=str).encode("utf-8") + + return StreamingResponse( + io.BytesIO(json_bytes), + media_type="application/json", + headers={"Content-Disposition": "attachment; filename=agents_tools_debug_export.json"}, + ) diff --git a/backend/pixelbot/routers/images.py b/backend/pixelbot/routers/images.py index 566f964..9be2da5 100644 --- a/backend/pixelbot/routers/images.py +++ b/backend/pixelbot/routers/images.py @@ -59,7 +59,6 @@ class GenerateImageRequest(BaseModel): @router.post("/generate_image", response_model=GenerateImageResponse) -@pxt_retry() def generate_image(body: GenerateImageRequest): """Generate an image using the configured provider (Gemini Imagen or OpenAI DALL-E). @@ -68,37 +67,29 @@ def generate_image(body: GenerateImageRequest): """ user_id = config.DEFAULT_USER_ID current_timestamp = datetime.now() + image_gen_table = pxt.get_table("pixelbot_v3.image_generation_tasks") + status = image_gen_table.insert( + [ImageGenRow(prompt=body.prompt, timestamp=current_timestamp, user_id=user_id)], + return_rows=True, + ) - try: - image_gen_table = pxt.get_table("pixelbot_v3.image_generation_tasks") - status = image_gen_table.insert( - [ImageGenRow(prompt=body.prompt, timestamp=current_timestamp, user_id=user_id)], - return_rows=True, - ) - - if not status.rows or status.rows[0].get("generated_image") is None: - raise HTTPException(status_code=500, detail="Image generation failed") - - img = status.rows[0]["generated_image"] - if not isinstance(img, Image.Image): - raise HTTPException(status_code=500, detail=f"Expected PIL Image, got {type(img)}") + if not status.rows or status.rows[0].get("generated_image") is None: + raise HTTPException(status_code=500, detail="Image generation failed") - buf = io.BytesIO() - img.save(buf, format="PNG") - img_base64 = base64.b64encode(buf.getvalue()).decode("utf-8") + img = status.rows[0]["generated_image"] + if not isinstance(img, Image.Image): + raise HTTPException(status_code=500, detail=f"Expected PIL Image, got {type(img)}") - return GenerateImageResponse( - generated_image_base64=img_base64, - timestamp=current_timestamp.isoformat(), - prompt=body.prompt, - provider="gemini", - ) + buf = io.BytesIO() + img.save(buf, format="PNG") + img_base64 = base64.b64encode(buf.getvalue()).decode("utf-8") - except HTTPException: - raise - except Exception as e: - logger.error(f"Error generating image: {e}", exc_info=True) - raise HTTPException(status_code=500, detail=str(e)) + return GenerateImageResponse( + generated_image_base64=img_base64, + timestamp=current_timestamp.isoformat(), + prompt=body.prompt, + provider="gemini", + ) # ── Image History ──────────────────────────────────────────────────────────── @@ -109,70 +100,63 @@ def generate_image(body: GenerateImageRequest): def get_image_history(): """Get history of generated images with provider metadata.""" user_id = config.DEFAULT_USER_ID + image_gen_table = pxt.get_table("pixelbot_v3.image_generation_tasks") - try: - image_gen_table = pxt.get_table("pixelbot_v3.image_generation_tasks") - - has_thumbnail_col = hasattr(image_gen_table, "thumbnail") - - select_kwargs: dict = { - "prompt": image_gen_table.prompt, - "timestamp": image_gen_table.timestamp, - "generated_image": image_gen_table.generated_image, - } - if has_thumbnail_col: - select_kwargs["thumbnail"] = image_gen_table.thumbnail - - results = ( - image_gen_table.where(image_gen_table.user_id == user_id) - .select(**select_kwargs) - .order_by(image_gen_table.timestamp, asc=False) - .limit(50) - .collect() - ) + has_thumbnail_col = hasattr(image_gen_table, "thumbnail") - image_history = [] - for entry in results: - img_data = entry.get("generated_image") - timestamp = entry.get("timestamp") - - if not isinstance(img_data, Image.Image): - continue - - thumbnail_b64 = entry.get("thumbnail") if has_thumbnail_col else None - if thumbnail_b64 and isinstance(thumbnail_b64, (str, bytes)): - if isinstance(thumbnail_b64, bytes): - thumbnail_b64 = thumbnail_b64.decode("utf-8") - if not thumbnail_b64.startswith("data:"): - thumbnail_b64 = f"data:image/png;base64,{thumbnail_b64}" - else: - thumbnail_b64 = create_thumbnail_base64(img_data, THUMB_SIZE) - - full_image_b64 = encode_image_base64(img_data) - - if thumbnail_b64 and full_image_b64: - image_history.append( - { - "prompt": entry.get("prompt"), - "timestamp": timestamp.strftime("%Y-%m-%d %H:%M:%S.%f") if timestamp else None, - "thumbnail_image": thumbnail_b64, - "full_image": full_image_b64, - "provider": "gemini", - } - ) + select_kwargs: dict = { + "prompt": image_gen_table.prompt, + "timestamp": image_gen_table.timestamp, + "generated_image": image_gen_table.generated_image, + } + if has_thumbnail_col: + select_kwargs["thumbnail"] = image_gen_table.thumbnail + + results = ( + image_gen_table.where(image_gen_table.user_id == user_id) + .select(**select_kwargs) + .order_by(image_gen_table.timestamp, asc=False) + .limit(50) + .collect() + ) + + image_history = [] + for entry in results: + img_data = entry.get("generated_image") + timestamp = entry.get("timestamp") + + if not isinstance(img_data, Image.Image): + continue + + thumbnail_b64 = entry.get("thumbnail") if has_thumbnail_col else None + if thumbnail_b64 and isinstance(thumbnail_b64, (str, bytes)): + if isinstance(thumbnail_b64, bytes): + thumbnail_b64 = thumbnail_b64.decode("utf-8") + if not thumbnail_b64.startswith("data:"): + thumbnail_b64 = f"data:image/png;base64,{thumbnail_b64}" + else: + thumbnail_b64 = create_thumbnail_base64(img_data, THUMB_SIZE) + + full_image_b64 = encode_image_base64(img_data) - return image_history + if thumbnail_b64 and full_image_b64: + image_history.append( + { + "prompt": entry.get("prompt"), + "timestamp": timestamp.strftime("%Y-%m-%d %H:%M:%S.%f") if timestamp else None, + "thumbnail_image": thumbnail_b64, + "full_image": full_image_b64, + "provider": "gemini", + } + ) - except Exception as e: - logger.error(f"Error fetching image history: {e}", exc_info=True) - raise HTTPException(status_code=500, detail=str(e)) + return image_history # ── Delete Generated Image ─────────────────────────────────────────────────── @router.delete("/delete_image/{timestamp_str}", response_model=DeleteResponse) -@pxt_retry() def delete_generated_image(timestamp_str: str): """Delete a generated image by timestamp.""" user_id = config.DEFAULT_USER_ID @@ -181,23 +165,15 @@ def delete_generated_image(timestamp_str: str): target_timestamp = datetime.strptime(timestamp_str, "%Y-%m-%d %H:%M:%S.%f") except ValueError: raise HTTPException(status_code=400, detail="Invalid timestamp format") + image_gen_table = pxt.get_table("pixelbot_v3.image_generation_tasks") + status = image_gen_table.delete( + where=(image_gen_table.timestamp == target_timestamp) & (image_gen_table.user_id == user_id) + ) - try: - image_gen_table = pxt.get_table("pixelbot_v3.image_generation_tasks") - status = image_gen_table.delete( - where=(image_gen_table.timestamp == target_timestamp) & (image_gen_table.user_id == user_id) - ) - - if status.num_rows == 0: - raise HTTPException(status_code=404, detail="No image found with that timestamp") + if status.num_rows == 0: + raise HTTPException(status_code=404, detail="No image found with that timestamp") - return DeleteResponse(message="Image deleted", num_deleted=status.num_rows) - - except HTTPException: - raise - except Exception as e: - logger.error(f"Error deleting image: {e}", exc_info=True) - raise HTTPException(status_code=500, detail=str(e)) + return DeleteResponse(message="Image deleted", num_deleted=status.num_rows) # ── Shared Request Models ──────────────────────────────────────────────────── @@ -217,7 +193,6 @@ class GenerateFluxImageRequest(BaseModel): @router.post("/generate_flux_image", response_model=GenerateImageResponse) -@pxt_retry() def generate_flux_image(body: GenerateFluxImageRequest): """Generate an image using BFL FLUX. @@ -228,45 +203,37 @@ def generate_flux_image(body: GenerateFluxImageRequest): w = max(64, (body.width // 16) * 16) h = max(64, (body.height // 16) * 16) + flux_table = pxt.get_table("pixelbot_v3.flux_generation_tasks") + status = flux_table.insert( + [ + FluxGenRow( + prompt=body.prompt, + width=w, + height=h, + timestamp=current_timestamp, + user_id=user_id, + ) + ], + return_rows=True, + ) - try: - flux_table = pxt.get_table("pixelbot_v3.flux_generation_tasks") - status = flux_table.insert( - [ - FluxGenRow( - prompt=body.prompt, - width=w, - height=h, - timestamp=current_timestamp, - user_id=user_id, - ) - ], - return_rows=True, - ) - - if not status.rows or status.rows[0].get("generated_image") is None: - raise HTTPException(status_code=500, detail="FLUX image generation failed") - - img = status.rows[0]["generated_image"] - if not isinstance(img, Image.Image): - raise HTTPException(status_code=500, detail=f"Expected PIL Image, got {type(img)}") + if not status.rows or status.rows[0].get("generated_image") is None: + raise HTTPException(status_code=500, detail="FLUX image generation failed") - buf = io.BytesIO() - img.save(buf, format="PNG") - img_base64 = base64.b64encode(buf.getvalue()).decode("utf-8") + img = status.rows[0]["generated_image"] + if not isinstance(img, Image.Image): + raise HTTPException(status_code=500, detail=f"Expected PIL Image, got {type(img)}") - return GenerateImageResponse( - generated_image_base64=img_base64, - timestamp=current_timestamp.isoformat(), - prompt=body.prompt, - provider="flux", - ) + buf = io.BytesIO() + img.save(buf, format="PNG") + img_base64 = base64.b64encode(buf.getvalue()).decode("utf-8") - except HTTPException: - raise - except Exception as e: - logger.error(f"Error generating FLUX image: {e}", exc_info=True) - raise HTTPException(status_code=500, detail=str(e)) + return GenerateImageResponse( + generated_image_base64=img_base64, + timestamp=current_timestamp.isoformat(), + prompt=body.prompt, + provider="flux", + ) @router.get("/flux_image_history") @@ -274,71 +241,64 @@ def generate_flux_image(body: GenerateFluxImageRequest): def get_flux_image_history(): """Get history of FLUX-generated images.""" user_id = config.DEFAULT_USER_ID + flux_table = pxt.get_table("pixelbot_v3.flux_generation_tasks") - try: - flux_table = pxt.get_table("pixelbot_v3.flux_generation_tasks") - - has_thumbnail_col = hasattr(flux_table, "thumbnail") - - select_kwargs: dict = { - "prompt": flux_table.prompt, - "timestamp": flux_table.timestamp, - "generated_image": flux_table.generated_image, - "width": flux_table.width, - "height": flux_table.height, - } - if has_thumbnail_col: - select_kwargs["thumbnail"] = flux_table.thumbnail - - results = ( - flux_table.where(flux_table.user_id == user_id) - .select(**select_kwargs) - .order_by(flux_table.timestamp, asc=False) - .limit(50) - .collect() - ) + has_thumbnail_col = hasattr(flux_table, "thumbnail") - history = [] - for entry in results: - img_data = entry.get("generated_image") - timestamp = entry.get("timestamp") - - if not isinstance(img_data, Image.Image): - continue - - thumbnail_b64 = entry.get("thumbnail") if has_thumbnail_col else None - if thumbnail_b64 and isinstance(thumbnail_b64, (str, bytes)): - if isinstance(thumbnail_b64, bytes): - thumbnail_b64 = thumbnail_b64.decode("utf-8") - if not thumbnail_b64.startswith("data:"): - thumbnail_b64 = f"data:image/png;base64,{thumbnail_b64}" - else: - thumbnail_b64 = create_thumbnail_base64(img_data, THUMB_SIZE) - - full_image_b64 = encode_image_base64(img_data) - - if thumbnail_b64 and full_image_b64: - history.append( - { - "prompt": entry.get("prompt"), - "timestamp": timestamp.strftime("%Y-%m-%d %H:%M:%S.%f") if timestamp else None, - "thumbnail_image": thumbnail_b64, - "full_image": full_image_b64, - "width": entry.get("width"), - "height": entry.get("height"), - "provider": "flux", - } - ) + select_kwargs: dict = { + "prompt": flux_table.prompt, + "timestamp": flux_table.timestamp, + "generated_image": flux_table.generated_image, + "width": flux_table.width, + "height": flux_table.height, + } + if has_thumbnail_col: + select_kwargs["thumbnail"] = flux_table.thumbnail + + results = ( + flux_table.where(flux_table.user_id == user_id) + .select(**select_kwargs) + .order_by(flux_table.timestamp, asc=False) + .limit(50) + .collect() + ) + + history = [] + for entry in results: + img_data = entry.get("generated_image") + timestamp = entry.get("timestamp") + + if not isinstance(img_data, Image.Image): + continue - return history + thumbnail_b64 = entry.get("thumbnail") if has_thumbnail_col else None + if thumbnail_b64 and isinstance(thumbnail_b64, (str, bytes)): + if isinstance(thumbnail_b64, bytes): + thumbnail_b64 = thumbnail_b64.decode("utf-8") + if not thumbnail_b64.startswith("data:"): + thumbnail_b64 = f"data:image/png;base64,{thumbnail_b64}" + else: + thumbnail_b64 = create_thumbnail_base64(img_data, THUMB_SIZE) - except Exception as e: - logger.error(f"Error fetching FLUX image history: {e}", exc_info=True) - raise HTTPException(status_code=500, detail=str(e)) + full_image_b64 = encode_image_base64(img_data) + + if thumbnail_b64 and full_image_b64: + history.append( + { + "prompt": entry.get("prompt"), + "timestamp": timestamp.strftime("%Y-%m-%d %H:%M:%S.%f") if timestamp else None, + "thumbnail_image": thumbnail_b64, + "full_image": full_image_b64, + "width": entry.get("width"), + "height": entry.get("height"), + "provider": "flux", + } + ) + + return history @router.post("/save_flux_image", response_model=SaveToCollectionResponse) -@pxt_retry() def save_flux_image_to_collection(body: SaveToCollectionRequest): """Save a FLUX-generated image into pixelbot_v3.images for CLIP embedding + RAG.""" user_id = config.DEFAULT_USER_ID @@ -347,49 +307,41 @@ def save_flux_image_to_collection(body: SaveToCollectionRequest): target_timestamp = datetime.strptime(body.timestamp, "%Y-%m-%d %H:%M:%S.%f") except ValueError: raise HTTPException(status_code=400, detail="Invalid timestamp format") + flux_table = pxt.get_table("pixelbot_v3.flux_generation_tasks") + result = ( + flux_table.where((flux_table.timestamp == target_timestamp) & (flux_table.user_id == user_id)) + .select(generated_image=flux_table.generated_image) + .collect() + ) - try: - flux_table = pxt.get_table("pixelbot_v3.flux_generation_tasks") - result = ( - flux_table.where((flux_table.timestamp == target_timestamp) & (flux_table.user_id == user_id)) - .select(generated_image=flux_table.generated_image) - .collect() - ) - - if len(result) == 0 or result[0].get("generated_image") is None: - raise HTTPException(status_code=404, detail="FLUX image not found") - - img = result[0]["generated_image"] - if not isinstance(img, Image.Image): - raise HTTPException(status_code=500, detail=f"Expected PIL Image, got {type(img)}") - - os.makedirs(config.UPLOAD_FOLDER, exist_ok=True) - file_uuid = str(uuid.uuid4()) - file_path = os.path.join(config.UPLOAD_FOLDER, f"{file_uuid}_flux.png") - img.save(file_path, format="PNG") + if len(result) == 0 or result[0].get("generated_image") is None: + raise HTTPException(status_code=404, detail="FLUX image not found") - images_table = pxt.get_table("pixelbot_v3.images") - images_table.insert( - [ - ImageRow( - image=file_path, - uuid=file_uuid, - timestamp=datetime.now(), - user_id=user_id, - ) - ] - ) + img = result[0]["generated_image"] + if not isinstance(img, Image.Image): + raise HTTPException(status_code=500, detail=f"Expected PIL Image, got {type(img)}") - return SaveToCollectionResponse( - message="FLUX image saved to collection — CLIP embedding and RAG indexing will run automatically", - uuid=file_uuid, - ) + os.makedirs(config.UPLOAD_FOLDER, exist_ok=True) + file_uuid = str(uuid.uuid4()) + file_path = os.path.join(config.UPLOAD_FOLDER, f"{file_uuid}_flux.png") + img.save(file_path, format="PNG") + + images_table = pxt.get_table("pixelbot_v3.images") + images_table.insert( + [ + ImageRow( + image=file_path, + uuid=file_uuid, + timestamp=datetime.now(), + user_id=user_id, + ) + ] + ) - except HTTPException: - raise - except Exception as e: - logger.error(f"Error saving FLUX image to collection: {e}", exc_info=True) - raise HTTPException(status_code=500, detail=str(e)) + return SaveToCollectionResponse( + message="FLUX image saved to collection — CLIP embedding and RAG indexing will run automatically", + uuid=file_uuid, + ) # ── Generate Video (Gemini Veo) ────────────────────────────────────────────── @@ -400,7 +352,6 @@ class GenerateVideoRequest(BaseModel): @router.post("/generate_video") -@pxt_retry() def generate_video(body: GenerateVideoRequest): """Generate a video using Gemini Veo. @@ -409,37 +360,29 @@ def generate_video(body: GenerateVideoRequest): """ user_id = config.DEFAULT_USER_ID current_timestamp = datetime.now() + video_gen_table = pxt.get_table("pixelbot_v3.video_generation_tasks") + status = video_gen_table.insert( + [VideoGenRow(prompt=body.prompt, timestamp=current_timestamp, user_id=user_id)], + return_rows=True, + ) - try: - video_gen_table = pxt.get_table("pixelbot_v3.video_generation_tasks") - status = video_gen_table.insert( - [VideoGenRow(prompt=body.prompt, timestamp=current_timestamp, user_id=user_id)], - return_rows=True, - ) - - if not status.rows or status.rows[0].get("generated_video") is None: - raise HTTPException(status_code=500, detail="Video generation failed") + if not status.rows or status.rows[0].get("generated_video") is None: + raise HTTPException(status_code=500, detail="Video generation failed") - video = status.rows[0]["generated_video"] + video = status.rows[0]["generated_video"] - # Pixeltable Video columns resolve to a file path string - video_path = str(video) if not isinstance(video, str) else video + # Pixeltable Video columns resolve to a file path string + video_path = str(video) if not isinstance(video, str) else video - if not os.path.exists(video_path): - raise HTTPException(status_code=500, detail="Generated video file not found on disk") + if not os.path.exists(video_path): + raise HTTPException(status_code=500, detail="Generated video file not found on disk") - return { - "timestamp": current_timestamp.isoformat(), - "prompt": body.prompt, - "provider": "gemini", - "video_path": video_path, - } - - except HTTPException: - raise - except Exception as e: - logger.error(f"Error generating video: {e}", exc_info=True) - raise HTTPException(status_code=500, detail=str(e)) + return { + "timestamp": current_timestamp.isoformat(), + "prompt": body.prompt, + "provider": "gemini", + "video_path": video_path, + } # ── Video History ──────────────────────────────────────────────────────────── @@ -450,45 +393,39 @@ def generate_video(body: GenerateVideoRequest): def get_video_history(): """Get history of generated videos.""" user_id = config.DEFAULT_USER_ID - - try: - video_gen_table = pxt.get_table("pixelbot_v3.video_generation_tasks") - - results = ( - video_gen_table.where(video_gen_table.user_id == user_id) - .select( - prompt=video_gen_table.prompt, - timestamp=video_gen_table.timestamp, - generated_video=video_gen_table.generated_video, - ) - .order_by(video_gen_table.timestamp, asc=False) - .limit(50) - .collect() + video_gen_table = pxt.get_table("pixelbot_v3.video_generation_tasks") + + results = ( + video_gen_table.where(video_gen_table.user_id == user_id) + .select( + prompt=video_gen_table.prompt, + timestamp=video_gen_table.timestamp, + generated_video=video_gen_table.generated_video, ) + .order_by(video_gen_table.timestamp, asc=False) + .limit(50) + .collect() + ) + + video_history = [] + for entry in results: + timestamp = entry.get("timestamp") + video = entry.get("generated_video") + + video_path = str(video) if video is not None else None + if not video_path or not os.path.exists(video_path): + continue - video_history = [] - for entry in results: - timestamp = entry.get("timestamp") - video = entry.get("generated_video") - - video_path = str(video) if video is not None else None - if not video_path or not os.path.exists(video_path): - continue - - video_history.append( - { - "prompt": entry.get("prompt"), - "timestamp": timestamp.strftime("%Y-%m-%d %H:%M:%S.%f") if timestamp else None, - "video_path": video_path, - "provider": "gemini", - } - ) - - return video_history + video_history.append( + { + "prompt": entry.get("prompt"), + "timestamp": timestamp.strftime("%Y-%m-%d %H:%M:%S.%f") if timestamp else None, + "video_path": video_path, + "provider": "gemini", + } + ) - except Exception as e: - logger.error(f"Error fetching video history: {e}", exc_info=True) - raise HTTPException(status_code=500, detail=str(e)) + return video_history # ── Serve Generated Video File ─────────────────────────────────────────────── @@ -505,7 +442,6 @@ def serve_generated_video(path: str): @router.delete("/delete_video/{timestamp_str}", response_model=DeleteResponse) -@pxt_retry() def delete_generated_video(timestamp_str: str): """Delete a generated video by timestamp.""" user_id = config.DEFAULT_USER_ID @@ -514,30 +450,21 @@ def delete_generated_video(timestamp_str: str): target_timestamp = datetime.strptime(timestamp_str, "%Y-%m-%d %H:%M:%S.%f") except ValueError: raise HTTPException(status_code=400, detail="Invalid timestamp format") + video_gen_table = pxt.get_table("pixelbot_v3.video_generation_tasks") + status = video_gen_table.delete( + where=(video_gen_table.timestamp == target_timestamp) & (video_gen_table.user_id == user_id) + ) - try: - video_gen_table = pxt.get_table("pixelbot_v3.video_generation_tasks") - status = video_gen_table.delete( - where=(video_gen_table.timestamp == target_timestamp) & (video_gen_table.user_id == user_id) - ) + if status.num_rows == 0: + raise HTTPException(status_code=404, detail="No video found with that timestamp") - if status.num_rows == 0: - raise HTTPException(status_code=404, detail="No video found with that timestamp") - - return DeleteResponse(message="Video deleted", num_deleted=status.num_rows) - - except HTTPException: - raise - except Exception as e: - logger.error(f"Error deleting video: {e}", exc_info=True) - raise HTTPException(status_code=500, detail=str(e)) + return DeleteResponse(message="Video deleted", num_deleted=status.num_rows) # ── Save Generated Media to Collection ─────────────────────────────────────── @router.post("/save_generated_image", response_model=SaveToCollectionResponse) -@pxt_retry() def save_generated_image_to_collection(body: SaveToCollectionRequest): """Save a generated image into pixelbot_v3.images so it enters the CLIP embedding + RAG pipeline.""" user_id = config.DEFAULT_USER_ID @@ -546,53 +473,44 @@ def save_generated_image_to_collection(body: SaveToCollectionRequest): target_timestamp = datetime.strptime(body.timestamp, "%Y-%m-%d %H:%M:%S.%f") except ValueError: raise HTTPException(status_code=400, detail="Invalid timestamp format") + gen_table = pxt.get_table("pixelbot_v3.image_generation_tasks") + result = ( + gen_table.where((gen_table.timestamp == target_timestamp) & (gen_table.user_id == user_id)) + .select(generated_image=gen_table.generated_image) + .collect() + ) - try: - gen_table = pxt.get_table("pixelbot_v3.image_generation_tasks") - result = ( - gen_table.where((gen_table.timestamp == target_timestamp) & (gen_table.user_id == user_id)) - .select(generated_image=gen_table.generated_image) - .collect() - ) - - if len(result) == 0 or result[0].get("generated_image") is None: - raise HTTPException(status_code=404, detail="Generated image not found") - - img = result[0]["generated_image"] - if not isinstance(img, Image.Image): - raise HTTPException(status_code=500, detail=f"Expected PIL Image, got {type(img)}") - - os.makedirs(config.UPLOAD_FOLDER, exist_ok=True) - file_uuid = str(uuid.uuid4()) - file_path = os.path.join(config.UPLOAD_FOLDER, f"{file_uuid}_generated.png") - img.save(file_path, format="PNG") + if len(result) == 0 or result[0].get("generated_image") is None: + raise HTTPException(status_code=404, detail="Generated image not found") - images_table = pxt.get_table("pixelbot_v3.images") - images_table.insert( - [ - ImageRow( - image=file_path, - uuid=file_uuid, - timestamp=datetime.now(), - user_id=user_id, - ) - ] - ) + img = result[0]["generated_image"] + if not isinstance(img, Image.Image): + raise HTTPException(status_code=500, detail=f"Expected PIL Image, got {type(img)}") - return SaveToCollectionResponse( - message="Image saved to collection — CLIP embedding and RAG indexing will run automatically", - uuid=file_uuid, - ) + os.makedirs(config.UPLOAD_FOLDER, exist_ok=True) + file_uuid = str(uuid.uuid4()) + file_path = os.path.join(config.UPLOAD_FOLDER, f"{file_uuid}_generated.png") + img.save(file_path, format="PNG") + + images_table = pxt.get_table("pixelbot_v3.images") + images_table.insert( + [ + ImageRow( + image=file_path, + uuid=file_uuid, + timestamp=datetime.now(), + user_id=user_id, + ) + ] + ) - except HTTPException: - raise - except Exception as e: - logger.error(f"Error saving generated image to collection: {e}", exc_info=True) - raise HTTPException(status_code=500, detail=str(e)) + return SaveToCollectionResponse( + message="Image saved to collection — CLIP embedding and RAG indexing will run automatically", + uuid=file_uuid, + ) @router.post("/save_generated_video", response_model=SaveToCollectionResponse) -@pxt_retry() def save_generated_video_to_collection(body: SaveToCollectionRequest): """Save a generated video into pixelbot_v3.videos so it enters keyframe/transcription/RAG pipeline.""" user_id = config.DEFAULT_USER_ID @@ -601,55 +519,46 @@ def save_generated_video_to_collection(body: SaveToCollectionRequest): target_timestamp = datetime.strptime(body.timestamp, "%Y-%m-%d %H:%M:%S.%f") except ValueError: raise HTTPException(status_code=400, detail="Invalid timestamp format") + gen_table = pxt.get_table("pixelbot_v3.video_generation_tasks") + result = ( + gen_table.where((gen_table.timestamp == target_timestamp) & (gen_table.user_id == user_id)) + .select(generated_video=gen_table.generated_video) + .collect() + ) + + if len(result) == 0 or result[0].get("generated_video") is None: + raise HTTPException(status_code=404, detail="Generated video not found") + + video = result[0]["generated_video"] + video_path = str(video) if not isinstance(video, str) else video + + if not os.path.exists(video_path): + raise HTTPException(status_code=500, detail="Generated video file not found on disk") + + file_uuid = str(uuid.uuid4()) + + videos_table = pxt.get_table("pixelbot_v3.videos") + videos_table.insert( + [ + VideoRow( + video=video_path, + uuid=file_uuid, + timestamp=datetime.now(), + user_id=user_id, + ) + ] + ) - try: - gen_table = pxt.get_table("pixelbot_v3.video_generation_tasks") - result = ( - gen_table.where((gen_table.timestamp == target_timestamp) & (gen_table.user_id == user_id)) - .select(generated_video=gen_table.generated_video) - .collect() - ) - - if len(result) == 0 or result[0].get("generated_video") is None: - raise HTTPException(status_code=404, detail="Generated video not found") - - video = result[0]["generated_video"] - video_path = str(video) if not isinstance(video, str) else video - - if not os.path.exists(video_path): - raise HTTPException(status_code=500, detail="Generated video file not found on disk") - - file_uuid = str(uuid.uuid4()) - - videos_table = pxt.get_table("pixelbot_v3.videos") - videos_table.insert( - [ - VideoRow( - video=video_path, - uuid=file_uuid, - timestamp=datetime.now(), - user_id=user_id, - ) - ] - ) - - return SaveToCollectionResponse( - message="Video saved to collection — keyframe extraction, transcription, and RAG indexing will run automatically", - uuid=file_uuid, - ) - - except HTTPException: - raise - except Exception as e: - logger.error(f"Error saving generated video to collection: {e}", exc_info=True) - raise HTTPException(status_code=500, detail=str(e)) + return SaveToCollectionResponse( + message="Video saved to collection — keyframe extraction, transcription, and RAG indexing will run automatically", + uuid=file_uuid, + ) # ── Generate Slideshow ─────────────────────────────────────────────────────── @router.post("/generate_slideshow", response_model=GenerateSlideshowResponse) -@pxt_retry() def generate_slideshow(body: GenerateSlideshowRequest): """Generate a video slideshow from a list of generated image timestamps.""" user_id = config.DEFAULT_USER_ID @@ -765,10 +674,6 @@ def generate_slideshow(body: GenerateSlideshowRequest): return GenerateSlideshowResponse( video_url=f"/api/serve_video?path={dest_path}", video_path=dest_path, uuid=file_uuid ) - - except Exception as e: - logger.error(f"Error generating slideshow: {e}", exc_info=True) - raise HTTPException(status_code=500, detail=str(e)) finally: try: pxt.drop_table(temp_table_name, force=True) @@ -787,47 +692,38 @@ class GenerateSpeechRequest(BaseModel): @router.post("/generate_speech", response_model=GenerateSpeechResponse) -@pxt_retry() def generate_speech(body: GenerateSpeechRequest): """Generate speech from text using OpenAI TTS via Pixeltable computed column.""" user_id = config.DEFAULT_USER_ID current_timestamp = datetime.now() voice = body.voice if body.voice in TTS_VOICES else "alloy" + speech_table = pxt.get_table("pixelbot_v3.speech_tasks") + status = speech_table.insert( + [ + SpeechTaskRow( + input_text=body.text, + voice=voice, + timestamp=current_timestamp, + user_id=user_id, + ) + ], + return_rows=True, + ) - try: - speech_table = pxt.get_table("pixelbot_v3.speech_tasks") - status = speech_table.insert( - [ - SpeechTaskRow( - input_text=body.text, - voice=voice, - timestamp=current_timestamp, - user_id=user_id, - ) - ], - return_rows=True, - ) - - if not status.rows or status.rows[0].get("audio") is None: - raise HTTPException(status_code=500, detail="Speech generation failed") - - audio_path = str(status.rows[0]["audio"]) - if not os.path.exists(audio_path): - raise HTTPException(status_code=500, detail="Audio file not found on disk") + if not status.rows or status.rows[0].get("audio") is None: + raise HTTPException(status_code=500, detail="Speech generation failed") - return GenerateSpeechResponse( - audio_url=f"/api/serve_audio?path={audio_path}", - audio_path=audio_path, - timestamp=current_timestamp.isoformat(), - voice=voice, - ) + audio_path = str(status.rows[0]["audio"]) + if not os.path.exists(audio_path): + raise HTTPException(status_code=500, detail="Audio file not found on disk") - except HTTPException: - raise - except Exception as e: - logger.error(f"Error generating speech: {e}", exc_info=True) - raise HTTPException(status_code=500, detail=str(e)) + return GenerateSpeechResponse( + audio_url=f"/api/serve_audio?path={audio_path}", + audio_path=audio_path, + timestamp=current_timestamp.isoformat(), + voice=voice, + ) class SaveSpeechRequest(BaseModel): @@ -835,38 +731,29 @@ class SaveSpeechRequest(BaseModel): @router.post("/save_generated_speech", response_model=SaveToCollectionResponse) -@pxt_retry() def save_generated_speech_to_collection(body: SaveSpeechRequest): """Save a TTS audio file into pixelbot_v3.audios so it enters the transcription + RAG pipeline.""" user_id = config.DEFAULT_USER_ID if not os.path.exists(body.audio_path): raise HTTPException(status_code=404, detail="Audio file not found on disk") + file_uuid = str(uuid.uuid4()) + audios_table = pxt.get_table("pixelbot_v3.audios") + audios_table.insert( + [ + AudioRow( + audio=body.audio_path, + uuid=file_uuid, + timestamp=datetime.now(), + user_id=user_id, + ) + ] + ) - try: - file_uuid = str(uuid.uuid4()) - audios_table = pxt.get_table("pixelbot_v3.audios") - audios_table.insert( - [ - AudioRow( - audio=body.audio_path, - uuid=file_uuid, - timestamp=datetime.now(), - user_id=user_id, - ) - ] - ) - - return SaveToCollectionResponse( - message="Audio saved to collection — transcription and RAG indexing will run automatically", - uuid=file_uuid, - ) - - except HTTPException: - raise - except Exception as e: - logger.error(f"Error saving speech to collection: {e}", exc_info=True) - raise HTTPException(status_code=500, detail=str(e)) + return SaveToCollectionResponse( + message="Audio saved to collection — transcription and RAG indexing will run automatically", + uuid=file_uuid, + ) @router.get("/serve_audio") @@ -874,16 +761,3 @@ def serve_audio(path: str): """Serve a generated audio file by path.""" safe_path = resolve_served_media_path(path) return FileResponse(safe_path, media_type="audio/wav", filename=os.path.basename(safe_path)) - - -@router.get("/tts_voices") -def get_tts_voices(): - """Return available TTS voice options.""" - return [ - {"id": "alloy", "label": "Alloy", "style": "Neutral, balanced"}, - {"id": "echo", "label": "Echo", "style": "Warm, conversational"}, - {"id": "fable", "label": "Fable", "style": "Expressive, storytelling"}, - {"id": "onyx", "label": "Onyx", "style": "Deep, authoritative"}, - {"id": "nova", "label": "Nova", "style": "Friendly, upbeat"}, - {"id": "shimmer", "label": "Shimmer", "style": "Clear, professional"}, - ] diff --git a/backend/pixelbot/routers/integrations.py b/backend/pixelbot/routers/integrations.py index 96b9f99..16ea30f 100644 --- a/backend/pixelbot/routers/integrations.py +++ b/backend/pixelbot/routers/integrations.py @@ -4,7 +4,6 @@ from datetime import datetime import pixeltable as pxt -import requests as http_requests from fastapi import APIRouter, Query from pixelbot import config @@ -17,6 +16,7 @@ TestNotificationRequest, TestNotificationResponse, ) +from pixelbot.notifications import deliver_notification, redacted_destination from pixelbot.utils import pxt_retry logger = logging.getLogger(__name__) @@ -38,14 +38,13 @@ def get_integrations_status(): @router.post("/test", response_model=TestNotificationResponse) -@pxt_retry() def test_notification(req: TestNotificationRequest): """Send a test notification and log it to the notifications table.""" service = req.service.lower() - now = datetime.utcnow() + now = datetime.now() - result = _send_notification(service, req.message) - if result is None: + delivery = deliver_notification(service, req.message) + if delivery is None: return TestNotificationResponse( service=service, status="error", @@ -53,23 +52,21 @@ def test_notification(req: TestNotificationRequest): timestamp=now.isoformat(), ) - is_success = "successfully" in result.lower() or "delivered" in result.lower() - notifications = pxt.get_table("pixelbot_v3.notifications") row = NotificationRow( service=service, - destination=_get_destination(service), + destination=redacted_destination(service), message=req.message, - status="success" if is_success else "error", - response_code=200 if is_success else 0, + status="success" if delivery.success else "error", + response_code=delivery.response_code, timestamp=now, ) notifications.insert([row]) return TestNotificationResponse( service=service, - status="success" if is_success else "error", - result=result, + status="success" if delivery.success else "error", + result=delivery.message, timestamp=now.isoformat(), ) @@ -152,59 +149,3 @@ def get_notification_log(limit: int = Query(default=50, ge=1, le=100)): entries.sort(key=lambda e: e.timestamp, reverse=True) return NotificationLogResponse(notifications=entries[:limit], total=len(entries)) - - -def _send_notification(service: str, message: str) -> str | None: - """Call the notification service directly (not via Pixeltable UDF).""" - try: - if service == "slack": - url = config.SLACK_WEBHOOK_URL - if not url: - return "Error: SLACK_WEBHOOK_URL not configured." - resp = http_requests.post(url, json={"text": message}, timeout=10) - return ( - "Slack message sent successfully." - if resp.status_code == 200 - else f"Slack error ({resp.status_code}): {resp.text}" - ) - - if service == "discord": - url = config.DISCORD_WEBHOOK_URL - if not url: - return "Error: DISCORD_WEBHOOK_URL not configured." - resp = http_requests.post(url, json={"content": message}, timeout=10) - return ( - "Discord message sent successfully." - if resp.status_code in (200, 204) - else f"Discord error ({resp.status_code}): {resp.text}" - ) - - if service == "webhook": - url = config.WEBHOOK_URL - if not url: - return "Error: WEBHOOK_URL not configured." - payload = {"text": message, "source": "pixelbot", "timestamp": datetime.utcnow().isoformat()} - resp = http_requests.post(url, json=payload, timeout=10) - return ( - f"Webhook delivered ({resp.status_code})." - if resp.status_code < 300 - else f"Webhook error ({resp.status_code}): {resp.text}" - ) - - return None - except http_requests.RequestException as e: - return f"{service} request failed: {e}" - - -def _get_destination(service: str) -> str: - url_map = { - "slack": config.SLACK_WEBHOOK_URL, - "discord": config.DISCORD_WEBHOOK_URL, - "webhook": config.WEBHOOK_URL, - } - url = url_map.get(service, "") - if not url: - return "(not configured)" - if len(url) > 40: - return url[:20] + "..." + url[-15:] - return url diff --git a/backend/pixelbot/routers/memory.py b/backend/pixelbot/routers/memory.py index fa32160..2916a60 100644 --- a/backend/pixelbot/routers/memory.py +++ b/backend/pixelbot/routers/memory.py @@ -7,7 +7,6 @@ from pixelbot import config from pixelbot.models import DeleteMemoryResponse, MemoryBankRow, MessageResponse -from pixelbot.utils import pxt_retry logger = logging.getLogger(__name__) router = APIRouter(prefix="/api", tags=["memory"]) @@ -53,7 +52,6 @@ def _insert_memory(body: SaveMemoryRequest) -> dict: @router.post("/memory", status_code=201, response_model=MessageResponse) -@pxt_retry() def save_memory(body: SaveMemoryRequest): """Save a memory item (code or text).""" return _insert_memory(body) @@ -63,7 +61,6 @@ def save_memory(body: SaveMemoryRequest): @router.delete("/memory/{timestamp_str}", response_model=DeleteMemoryResponse) -@pxt_retry() def delete_memory(timestamp_str: str): """Delete a memory item by timestamp.""" user_id = config.DEFAULT_USER_ID @@ -72,20 +69,10 @@ def delete_memory(timestamp_str: str): target_timestamp = datetime.strptime(timestamp_str, "%Y-%m-%d %H:%M:%S.%f") except ValueError: raise HTTPException(status_code=400, detail="Invalid timestamp format") + memory_table = pxt.get_table("pixelbot_v3.memory_bank") + status = memory_table.delete(where=(memory_table.timestamp == target_timestamp) & (memory_table.user_id == user_id)) - try: - memory_table = pxt.get_table("pixelbot_v3.memory_bank") - status = memory_table.delete( - where=(memory_table.timestamp == target_timestamp) & (memory_table.user_id == user_id) - ) - - if status.num_rows == 0: - raise HTTPException(status_code=404, detail="No memory item found with that timestamp") - - return DeleteMemoryResponse(message="Memory item deleted", num_deleted=status.num_rows) + if status.num_rows == 0: + raise HTTPException(status_code=404, detail="No memory item found with that timestamp") - except HTTPException: - raise - except Exception as e: - logger.error(f"Error deleting memory: {e}", exc_info=True) - raise HTTPException(status_code=500, detail=str(e)) + return DeleteMemoryResponse(message="Memory item deleted", num_deleted=status.num_rows) diff --git a/backend/pixelbot/routers/personas.py b/backend/pixelbot/routers/personas.py index 8cccaeb..50d75a7 100644 --- a/backend/pixelbot/routers/personas.py +++ b/backend/pixelbot/routers/personas.py @@ -7,7 +7,6 @@ from pixelbot import config from pixelbot.models import DeleteResponse, MessageResponse, UserPersonaRow -from pixelbot.utils import pxt_retry logger = logging.getLogger(__name__) router = APIRouter(prefix="/api", tags=["personas"]) @@ -30,106 +29,79 @@ class PersonaUpdateRequest(BaseModel): @router.post("/personas", status_code=201, response_model=MessageResponse) -@pxt_retry() def create_persona(body: PersonaRequest): """Create a new persona.""" user_id = config.DEFAULT_USER_ID if not body.persona_name.strip(): raise HTTPException(status_code=400, detail="Persona name cannot be empty") + personas_table = pxt.get_table("pixelbot_v3.user_personas") try: - personas_table = pxt.get_table("pixelbot_v3.user_personas") - - try: - personas_table.insert( - [ - UserPersonaRow( - user_id=user_id, - persona_name=body.persona_name.strip(), - initial_prompt=body.initial_prompt, - final_prompt=body.final_prompt, - llm_params=body.llm_params, - timestamp=datetime.now(), - ) - ] - ) - return {"message": f"Persona '{body.persona_name}' created successfully"} - - except Exception as insert_err: - err_str = str(insert_err).lower() - if "unique constraint" in err_str or "primary key constraint" in err_str: - raise HTTPException(status_code=409, detail=f"Persona '{body.persona_name}' already exists") - raise insert_err - - except HTTPException: - raise - except Exception as e: - logger.error(f"Error creating persona: {e}", exc_info=True) - raise HTTPException(status_code=500, detail=str(e)) + personas_table.insert( + [ + UserPersonaRow( + user_id=user_id, + persona_name=body.persona_name.strip(), + initial_prompt=body.initial_prompt, + final_prompt=body.final_prompt, + llm_params=body.llm_params, + timestamp=datetime.now(), + ) + ] + ) + return {"message": f"Persona '{body.persona_name}' created successfully"} + + except Exception as insert_err: + err_str = str(insert_err).lower() + if "unique constraint" in err_str or "primary key constraint" in err_str: + raise HTTPException(status_code=409, detail=f"Persona '{body.persona_name}' already exists") + raise insert_err # ── Update Persona ──────────────────────────────────────────────────────────── @router.put("/personas/{persona_name:path}", response_model=MessageResponse) -@pxt_retry() def update_persona(persona_name: str, body: PersonaUpdateRequest): """Update an existing persona.""" user_id = config.DEFAULT_USER_ID if not persona_name: raise HTTPException(status_code=400, detail="Persona name is required") + personas_table = pxt.get_table("pixelbot_v3.user_personas") + status = personas_table.update( + { + "initial_prompt": body.initial_prompt, + "final_prompt": body.final_prompt, + "llm_params": body.llm_params, + "timestamp": datetime.now(), + }, + where=(personas_table.user_id == user_id) & (personas_table.persona_name == persona_name), + ) - try: - personas_table = pxt.get_table("pixelbot_v3.user_personas") - status = personas_table.update( - { - "initial_prompt": body.initial_prompt, - "final_prompt": body.final_prompt, - "llm_params": body.llm_params, - "timestamp": datetime.now(), - }, - where=(personas_table.user_id == user_id) & (personas_table.persona_name == persona_name), - ) - - if status.num_rows == 0: - raise HTTPException(status_code=404, detail="Persona not found") + if status.num_rows == 0: + raise HTTPException(status_code=404, detail="Persona not found") - return {"message": f"Persona '{persona_name}' updated successfully"} - - except HTTPException: - raise - except Exception as e: - logger.error(f"Error updating persona: {e}", exc_info=True) - raise HTTPException(status_code=500, detail=str(e)) + return {"message": f"Persona '{persona_name}' updated successfully"} # ── Delete Persona ──────────────────────────────────────────────────────────── @router.delete("/personas/{persona_name:path}", response_model=DeleteResponse) -@pxt_retry() def delete_persona(persona_name: str): """Delete a persona by name.""" user_id = config.DEFAULT_USER_ID if not persona_name: raise HTTPException(status_code=400, detail="Persona name is required") + personas_table = pxt.get_table("pixelbot_v3.user_personas") + status = personas_table.delete( + where=(personas_table.user_id == user_id) & (personas_table.persona_name == persona_name) + ) - try: - personas_table = pxt.get_table("pixelbot_v3.user_personas") - status = personas_table.delete( - where=(personas_table.user_id == user_id) & (personas_table.persona_name == persona_name) - ) - - if status.num_rows == 0: - raise HTTPException(status_code=404, detail="Persona not found") - - return {"message": f"Persona '{persona_name}' deleted successfully", "num_deleted": status.num_rows} + if status.num_rows == 0: + raise HTTPException(status_code=404, detail="Persona not found") - except HTTPException: - raise - except Exception as e: - logger.error(f"Error deleting persona: {e}", exc_info=True) - raise HTTPException(status_code=500, detail=str(e)) + return {"message": f"Persona '{persona_name}' deleted successfully", "num_deleted": status.num_rows} diff --git a/backend/pixelbot/routers/studio.py b/backend/pixelbot/routers/studio.py index 8d5e255..c7e8534 100644 --- a/backend/pixelbot/routers/studio.py +++ b/backend/pixelbot/routers/studio.py @@ -525,80 +525,64 @@ class CsvRowsRequest(BaseModel): def get_csv_rows(body: CsvRowsRequest): """Return paginated rows from a CSV table.""" user_id = config.DEFAULT_USER_ID + registration = _verify_csv_ownership(body.csv_uuid, user_id) + table_name = registration["table_name"] + col_names = registration["col_names"] + total_rows = registration["row_count"] - try: - registration = _verify_csv_ownership(body.csv_uuid, user_id) - table_name = registration["table_name"] - col_names = registration["col_names"] - total_rows = registration["row_count"] - - # Fetch rows from the actual CSV table - tbl = pxt.get_table(table_name) - rows_data: list[dict] = [] - all_rows = list(tbl.select().limit(body.limit + body.offset).collect()) - - for row in all_rows[body.offset :]: - row_dict: dict = {} - for col in col_names: - val = row.get(col) - # Ensure JSON-serializable values - if val is None: - row_dict[col] = None - elif isinstance(val, (int, float, bool, str)): - row_dict[col] = val - else: - row_dict[col] = str(val) - rows_data.append(row_dict) - - return { - "table_name": table_name, - "columns": col_names, - "rows": rows_data, - "total": total_rows, - "offset": body.offset, - "limit": body.limit, - } + # Fetch rows from the actual CSV table + tbl = pxt.get_table(table_name) + rows_data: list[dict] = [] + all_rows = list(tbl.select().limit(body.limit + body.offset).collect()) - except HTTPException: - raise - except Exception as e: - logger.error(f"Error fetching CSV rows: {e}", exc_info=True) - raise HTTPException(status_code=500, detail=str(e)) + for row in all_rows[body.offset :]: + row_dict: dict = {} + for col in col_names: + val = row.get(col) + # Ensure JSON-serializable values + if val is None: + row_dict[col] = None + elif isinstance(val, (int, float, bool, str)): + row_dict[col] = val + else: + row_dict[col] = str(val) + rows_data.append(row_dict) + + return { + "table_name": table_name, + "columns": col_names, + "rows": rows_data, + "total": total_rows, + "offset": body.offset, + "limit": body.limit, + } @router.delete("/csv/{csv_uuid}") def delete_csv_table(csv_uuid: str): """Delete a CSV table and its registry entry.""" user_id = config.DEFAULT_USER_ID + registry = pxt.get_table("pixelbot_v3.csv_registry") + check = ( + registry.where((registry.uuid == csv_uuid) & (registry.user_id == user_id)) + .select(registry.table_name) + .collect() + ) + if not check: + raise HTTPException(status_code=404, detail="CSV table not found") - try: - registry = pxt.get_table("pixelbot_v3.csv_registry") - check = ( - registry.where((registry.uuid == csv_uuid) & (registry.user_id == user_id)) - .select(registry.table_name) - .collect() - ) - if not check: - raise HTTPException(status_code=404, detail="CSV table not found") - - table_name = check[0]["table_name"] - - # Drop the actual CSV table - try: - pxt.drop_table(table_name, force=True) - except Exception as e: - logger.warning(f"Could not drop CSV table {table_name}: {e}") + table_name = check[0]["table_name"] - # Remove from registry - registry.delete(where=(registry.uuid == csv_uuid) & (registry.user_id == user_id)) + # Drop the actual CSV table + try: + pxt.drop_table(table_name, force=True) + except Exception as e: + logger.warning(f"Could not drop CSV table {table_name}: {e}") - return {"message": f"CSV table '{table_name}' deleted"} + # Remove from registry + registry.delete(where=(registry.uuid == csv_uuid) & (registry.user_id == user_id)) - except HTTPException: - raise - except Exception as e: - logger.error(f"Error deleting CSV table: {e}", exc_info=True) - raise HTTPException(status_code=500, detail=str(e)) + return {"message": f"CSV table '{table_name}' deleted"} # ── CSV Row CRUD ───────────────────────────────────────────────────────────── @@ -699,32 +683,24 @@ def csv_add_rows(body: CsvAddRowsRequest): if not body.rows: raise HTTPException(status_code=400, detail="No rows provided") + tbl = pxt.get_table(table_name) + schema = _get_col_schema(tbl) - try: - tbl = pxt.get_table(table_name) - schema = _get_col_schema(tbl) + coerced_rows = [] + for row in body.rows: + coerced = {} + for col in col_names: + coerced[col] = _coerce_value(row.get(col), col, schema) + coerced_rows.append(coerced) - coerced_rows = [] - for row in body.rows: - coerced = {} - for col in col_names: - coerced[col] = _coerce_value(row.get(col), col, schema) - coerced_rows.append(coerced) + tbl.insert(coerced_rows) + new_count = _sync_registry_row_count(table_name, user_id) - tbl.insert(coerced_rows) - new_count = _sync_registry_row_count(table_name, user_id) - - return { - "message": f"Added {len(coerced_rows)} row(s)", - "rows_added": len(coerced_rows), - "new_total": new_count, - } - - except HTTPException: - raise - except Exception as e: - logger.error(f"Error adding CSV rows: {e}", exc_info=True) - raise HTTPException(status_code=500, detail=str(e)) + return { + "message": f"Added {len(coerced_rows)} row(s)", + "rows_added": len(coerced_rows), + "new_total": new_count, + } class CsvUpdateRowRequest(BaseModel): @@ -747,39 +723,31 @@ def csv_update_row(body: CsvUpdateRowRequest): if not body.updated_values: raise HTTPException(status_code=400, detail="No updated values provided") + tbl = pxt.get_table(table_name) + schema = _get_col_schema(tbl) - try: - tbl = pxt.get_table(table_name) - schema = _get_col_schema(tbl) - - # Coerce original row values - coerced_original: dict = {} - for col in col_names: - coerced_original[col] = _coerce_value(body.original_row.get(col), col, schema) - - where = _build_row_where(tbl, col_names, coerced_original) + # Coerce original row values + coerced_original: dict = {} + for col in col_names: + coerced_original[col] = _coerce_value(body.original_row.get(col), col, schema) - # Coerce updated values - update_dict: dict = {} - for col, val in body.updated_values.items(): - if col in col_names: - update_dict[col] = _coerce_value(val, col, schema) + where = _build_row_where(tbl, col_names, coerced_original) - if not update_dict: - raise HTTPException(status_code=400, detail="No valid columns to update") + # Coerce updated values + update_dict: dict = {} + for col, val in body.updated_values.items(): + if col in col_names: + update_dict[col] = _coerce_value(val, col, schema) - status = tbl.update(update_dict, where=where) + if not update_dict: + raise HTTPException(status_code=400, detail="No valid columns to update") - return { - "message": f"Updated {status.num_rows} row(s)", - "rows_updated": status.num_rows, - } + status = tbl.update(update_dict, where=where) - except HTTPException: - raise - except Exception as e: - logger.error(f"Error updating CSV row: {e}", exc_info=True) - raise HTTPException(status_code=500, detail=str(e)) + return { + "message": f"Updated {status.num_rows} row(s)", + "rows_updated": status.num_rows, + } class CsvDeleteRowsRequest(BaseModel): @@ -794,30 +762,22 @@ def csv_delete_rows(body: CsvDeleteRowsRequest): reg = _verify_csv_ownership(body.csv_uuid, user_id) table_name = reg["table_name"] col_names: list[str] = reg["col_names"] + tbl = pxt.get_table(table_name) + schema = _get_col_schema(tbl) - try: - tbl = pxt.get_table(table_name) - schema = _get_col_schema(tbl) - - coerced: dict = {} - for col in col_names: - coerced[col] = _coerce_value(body.row_values.get(col), col, schema) - - where = _build_row_where(tbl, col_names, coerced) - status = tbl.delete(where=where) - new_count = _sync_registry_row_count(table_name, user_id) + coerced: dict = {} + for col in col_names: + coerced[col] = _coerce_value(body.row_values.get(col), col, schema) - return { - "message": f"Deleted {status.num_rows} row(s)", - "rows_deleted": status.num_rows, - "new_total": new_count, - } + where = _build_row_where(tbl, col_names, coerced) + status = tbl.delete(where=where) + new_count = _sync_registry_row_count(table_name, user_id) - except HTTPException: - raise - except Exception as e: - logger.error(f"Error deleting CSV rows: {e}", exc_info=True) - raise HTTPException(status_code=500, detail=str(e)) + return { + "message": f"Deleted {status.num_rows} row(s)", + "rows_deleted": status.num_rows, + "new_total": new_count, + } class CsvRevertRequest(BaseModel): @@ -830,25 +790,19 @@ def csv_revert(body: CsvRevertRequest): user_id = config.DEFAULT_USER_ID registration = _verify_csv_ownership(body.csv_uuid, user_id) table_name = registration["table_name"] + tbl = pxt.get_table(table_name) + tbl.revert() + new_count = _sync_registry_row_count(table_name, user_id) - try: - tbl = pxt.get_table(table_name) - tbl.revert() - new_count = _sync_registry_row_count(table_name, user_id) - - versions = tbl.get_versions() - can_undo = len(versions) > 1 + versions = tbl.get_versions() + can_undo = len(versions) > 1 - return { - "message": "Reverted to previous version", - "new_total": new_count, - "current_version": versions[0]["version"] if versions else 0, - "can_undo": can_undo, - } - - except Exception as e: - logger.error(f"Error reverting CSV table: {e}", exc_info=True) - raise HTTPException(status_code=500, detail=str(e)) + return { + "message": "Reverted to previous version", + "new_total": new_count, + "current_version": versions[0]["version"] if versions else 0, + "can_undo": can_undo, + } # ── CSV Version History ────────────────────────────────────────────────────── @@ -860,33 +814,27 @@ def csv_versions(csv_uuid: str): user_id = config.DEFAULT_USER_ID registration = _verify_csv_ownership(csv_uuid, user_id) table_name = registration["table_name"] - - try: - tbl = pxt.get_table(table_name) - versions = tbl.get_versions() - - return { - "table_name": table_name, - "current_version": versions[0]["version"] if versions else 0, - "can_undo": len(versions) > 1, - "versions": [ - { - "version": v["version"], - "created_at": v["created_at"].isoformat() if v.get("created_at") else None, - "change_type": v.get("change_type", "data"), - "inserts": v.get("inserts", 0), - "updates": v.get("updates", 0), - "deletes": v.get("deletes", 0), - "errors": v.get("errors", 0), - "schema_change": v.get("schema_change"), - } - for v in versions - ], - } - - except Exception as e: - logger.error(f"Error fetching CSV versions: {e}", exc_info=True) - raise HTTPException(status_code=500, detail=str(e)) + tbl = pxt.get_table(table_name) + versions = tbl.get_versions() + + return { + "table_name": table_name, + "current_version": versions[0]["version"] if versions else 0, + "can_undo": len(versions) > 1, + "versions": [ + { + "version": v["version"], + "created_at": v["created_at"].isoformat() if v.get("created_at") else None, + "change_type": v.get("change_type", "data"), + "inserts": v.get("inserts", 0), + "updates": v.get("updates", 0), + "deletes": v.get("deletes", 0), + "errors": v.get("errors", 0), + "schema_change": v.get("schema_change"), + } + for v in versions + ], + } # ── Cross-modal Similarity Search ──────────────────────────────────────────── @@ -900,7 +848,6 @@ class SearchRequest(BaseModel): @router.post("/search") -@pxt_retry() def search_studio(body: SearchRequest): """Cross-modal semantic search across all file types using embedding indexes.""" user_id = config.DEFAULT_USER_ID @@ -1097,7 +1044,6 @@ def _get_embed_clip_fn(): @router.get("/embeddings") -@pxt_retry() def get_embeddings(space: str = "text", limit: int = 200): """ Return 2-D UMAP-projected embeddings for visualization. @@ -1322,34 +1268,26 @@ def _collect_visual_embeddings( def get_image_preview(uuid: str): """Get a larger preview of an image for the studio workspace.""" user_id = config.DEFAULT_USER_ID - try: - img_table = _get_pxt_table("image") - rows = ( - img_table.where((img_table.uuid == uuid) & (img_table.user_id == user_id)) - .select(img=img_table.image) - .collect() - ) + img_table = _get_pxt_table("image") + rows = ( + img_table.where((img_table.uuid == uuid) & (img_table.user_id == user_id)).select(img=img_table.image).collect() + ) - if len(rows) == 0: - raise HTTPException(status_code=404, detail="Image not found") + if len(rows) == 0: + raise HTTPException(status_code=404, detail="Image not found") - img = rows[0]["img"] - if not isinstance(img, Image.Image): - raise HTTPException(status_code=500, detail="Could not load image") + img = rows[0]["img"] + if not isinstance(img, Image.Image): + raise HTTPException(status_code=500, detail="Could not load image") - width, height = img.size - preview = _pil_image_to_data_uri(img, max_size=PREVIEW_SIZE) - return { - "preview": preview, - "width": width, - "height": height, - "mode": img.mode, - } - except HTTPException: - raise - except Exception as e: - logger.error(f"Studio: error getting image preview: {e}", exc_info=True) - raise HTTPException(status_code=500, detail=str(e)) + width, height = img.size + preview = _pil_image_to_data_uri(img, max_size=PREVIEW_SIZE) + return { + "preview": preview, + "width": width, + "height": height, + "mode": img.mode, + } # ── Image Transform ────────────────────────────────────────────────────────── @@ -1366,37 +1304,30 @@ class TransformRequest(BaseModel): def transform_image(body: TransformRequest): """Apply a PIL transform to an image and return the preview (no storage).""" user_id = config.DEFAULT_USER_ID - try: - img_table = _get_pxt_table("image") - rows = ( - img_table.where((img_table.uuid == body.uuid) & (img_table.user_id == user_id)) - .select(img=img_table.image) - .collect() - ) - - if len(rows) == 0: - raise HTTPException(status_code=404, detail="Image not found") + img_table = _get_pxt_table("image") + rows = ( + img_table.where((img_table.uuid == body.uuid) & (img_table.user_id == user_id)) + .select(img=img_table.image) + .collect() + ) - img = rows[0]["img"] - if not isinstance(img, Image.Image): - raise HTTPException(status_code=500, detail="Could not load image") + if len(rows) == 0: + raise HTTPException(status_code=404, detail="Image not found") - result = _apply_image_operation(img, body.operation, body.params) - preview = _pil_image_to_data_uri(result, max_size=PREVIEW_SIZE) + img = rows[0]["img"] + if not isinstance(img, Image.Image): + raise HTTPException(status_code=500, detail="Could not load image") - return { - "preview": preview, - "width": result.size[0], - "height": result.size[1], - "mode": result.mode, - "operation": body.operation, - } + result = _apply_image_operation(img, body.operation, body.params) + preview = _pil_image_to_data_uri(result, max_size=PREVIEW_SIZE) - except HTTPException: - raise - except Exception as e: - logger.error(f"Studio: image transform error: {e}", exc_info=True) - raise HTTPException(status_code=500, detail=str(e)) + return { + "preview": preview, + "width": result.size[0], + "height": result.size[1], + "mode": result.mode, + "operation": body.operation, + } def _get_full_res_transform(body: TransformRequest) -> Image.Image: @@ -1515,235 +1446,211 @@ def detect_objects(body: DetectRequest): model_info = DETECTION_MODELS.get(body.model) if not model_info: raise HTTPException(status_code=400, detail=f"Unknown model: {body.model}") + img: Image.Image | None = None - # Load image from the right table - try: - img: Image.Image | None = None - - if body.source == "video_frame": - if body.frame_idx is None: - raise HTTPException(status_code=400, detail="frame_idx required for video_frame source") - frames_view = pxt.get_table("pixelbot_v3.video_frames") - rows = ( - frames_view.where( - (frames_view.uuid == body.uuid) - & (frames_view.user_id == user_id) - & (frames_view.frame_idx == body.frame_idx) - ) - .select(frame=frames_view.frame) - .collect() - ) - if rows: - img = rows[0].get("frame") - else: - img_table = _get_pxt_table("image") - rows = ( - img_table.where((img_table.uuid == body.uuid) & (img_table.user_id == user_id)) - .select(img=img_table.image) - .collect() + if body.source == "video_frame": + if body.frame_idx is None: + raise HTTPException(status_code=400, detail="frame_idx required for video_frame source") + frames_view = pxt.get_table("pixelbot_v3.video_frames") + rows = ( + frames_view.where( + (frames_view.uuid == body.uuid) + & (frames_view.user_id == user_id) + & (frames_view.frame_idx == body.frame_idx) ) - if rows: - img = rows[0].get("img") - - if img is None or not isinstance(img, Image.Image): - raise HTTPException(status_code=404, detail="Image not found") - - # Convert to RGB if needed - if img.mode != "RGB": - img = img.convert("RGB") - - except HTTPException: - raise - except Exception as e: - logger.error(f"Detection: error loading image: {e}", exc_info=True) - raise HTTPException(status_code=500, detail=f"Failed to load image: {e}") - - # Run inference - try: - processor, model = _get_detection_model(body.model) - img_width, img_height = img.size - - if model_info["type"] == "detection": - inputs = processor(images=img, return_tensors="pt") - with torch.no_grad(): - outputs = model(**inputs) - - target_sizes = torch.tensor([[img_height, img_width]]) - results = processor.post_process_object_detection( - outputs, target_sizes=target_sizes, threshold=body.threshold - )[0] - - detections = [] - for score, label_id, box in zip( - results["scores"].tolist(), - results["labels"].tolist(), - results["boxes"].tolist(), - ): - detections.append( - { - "label": model.config.id2label[label_id], - "score": round(score, 3), - "box": { - "x1": round(box[0], 1), - "y1": round(box[1], 1), - "x2": round(box[2], 1), - "y2": round(box[3], 1), - }, - } - ) + .select(frame=frames_view.frame) + .collect() + ) + if rows: + img = rows[0].get("frame") + else: + img_table = _get_pxt_table("image") + rows = ( + img_table.where((img_table.uuid == body.uuid) & (img_table.user_id == user_id)) + .select(img=img_table.image) + .collect() + ) + if rows: + img = rows[0].get("img") - # Sort by score descending - detections.sort(key=lambda d: d["score"], reverse=True) + if img is None or not isinstance(img, Image.Image): + raise HTTPException(status_code=404, detail="Image not found") - return { - "type": "detection", - "model": body.model, - "image_width": img_width, - "image_height": img_height, - "count": len(detections), - "detections": detections, - } + # Convert to RGB if needed + if img.mode != "RGB": + img = img.convert("RGB") + processor, model = _get_detection_model(body.model) + img_width, img_height = img.size + + if model_info["type"] == "detection": + inputs = processor(images=img, return_tensors="pt") + with torch.no_grad(): + outputs = model(**inputs) + + target_sizes = torch.tensor([[img_height, img_width]]) + results = processor.post_process_object_detection(outputs, target_sizes=target_sizes, threshold=body.threshold)[ + 0 + ] + + detections = [] + for score, label_id, box in zip( + results["scores"].tolist(), + results["labels"].tolist(), + results["boxes"].tolist(), + ): + detections.append( + { + "label": model.config.id2label[label_id], + "score": round(score, 3), + "box": { + "x1": round(box[0], 1), + "y1": round(box[1], 1), + "x2": round(box[2], 1), + "y2": round(box[3], 1), + }, + } + ) - elif model_info["type"] == "segmentation": - inputs = processor(images=img, return_tensors="pt") - with torch.no_grad(): - outputs = model(**inputs) - - result = processor.post_process_panoptic_segmentation( - outputs, threshold=body.threshold, target_sizes=[(img_height, img_width)] - )[0] - - seg_array = result["segmentation"].cpu().numpy() - segments = [] - for seg_info in result.get("segments_info", []): - seg_id = seg_info["id"] - label_id = seg_info["label_id"] - label_text = model.config.id2label.get(label_id, f"class_{label_id}") - score = round(seg_info.get("score", 0.0), 3) - - # Compute bounding box from segment mask - mask = seg_array == seg_id - ys, xs = mask.nonzero() - if len(ys) == 0: - continue + # Sort by score descending + detections.sort(key=lambda d: d["score"], reverse=True) - segments.append( - { - "id": int(seg_id), - "label": label_text, - "score": score, - "is_thing": seg_info.get("isthing", True), - "box": { - "x1": round(float(xs.min()), 1), - "y1": round(float(ys.min()), 1), - "x2": round(float(xs.max()), 1), - "y2": round(float(ys.max()), 1), - }, - "pixel_count": int(mask.sum()), - } - ) + return { + "type": "detection", + "model": body.model, + "image_width": img_width, + "image_height": img_height, + "count": len(detections), + "detections": detections, + } - segments.sort(key=lambda s: s["score"], reverse=True) + elif model_info["type"] == "segmentation": + inputs = processor(images=img, return_tensors="pt") + with torch.no_grad(): + outputs = model(**inputs) + + result = processor.post_process_panoptic_segmentation( + outputs, threshold=body.threshold, target_sizes=[(img_height, img_width)] + )[0] + + seg_array = result["segmentation"].cpu().numpy() + segments = [] + for seg_info in result.get("segments_info", []): + seg_id = seg_info["id"] + label_id = seg_info["label_id"] + label_text = model.config.id2label.get(label_id, f"class_{label_id}") + score = round(seg_info.get("score", 0.0), 3) + + # Compute bounding box from segment mask + mask = seg_array == seg_id + ys, xs = mask.nonzero() + if len(ys) == 0: + continue + + segments.append( + { + "id": int(seg_id), + "label": label_text, + "score": score, + "is_thing": seg_info.get("isthing", True), + "box": { + "x1": round(float(xs.min()), 1), + "y1": round(float(ys.min()), 1), + "x2": round(float(xs.max()), 1), + "y2": round(float(ys.max()), 1), + }, + "pixel_count": int(mask.sum()), + } + ) - return { - "type": "segmentation", - "model": body.model, - "image_width": img_width, - "image_height": img_height, - "count": len(segments), - "segments": segments, - } + segments.sort(key=lambda s: s["score"], reverse=True) - else: - inputs = processor(images=img, return_tensors="pt") - with torch.no_grad(): - outputs = model(**inputs) - - logits = outputs.logits[0] - probs = torch.nn.functional.softmax(logits, dim=-1) - top_k = min(body.top_k, len(probs)) - top_probs, top_indices = torch.topk(probs, top_k) - - classifications = [] - for prob, idx in zip(top_probs.tolist(), top_indices.tolist()): - classifications.append( - { - "label": model.config.id2label[idx], - "score": round(prob, 4), - } - ) + return { + "type": "segmentation", + "model": body.model, + "image_width": img_width, + "image_height": img_height, + "count": len(segments), + "segments": segments, + } - return { - "type": "classification", - "model": body.model, - "image_width": img_width, - "image_height": img_height, - "count": len(classifications), - "classifications": classifications, - } + else: + inputs = processor(images=img, return_tensors="pt") + with torch.no_grad(): + outputs = model(**inputs) + + logits = outputs.logits[0] + probs = torch.nn.functional.softmax(logits, dim=-1) + top_k = min(body.top_k, len(probs)) + top_probs, top_indices = torch.topk(probs, top_k) + + classifications = [] + for prob, idx in zip(top_probs.tolist(), top_indices.tolist()): + classifications.append( + { + "label": model.config.id2label[idx], + "score": round(prob, 4), + } + ) - except Exception as e: - logger.error(f"Detection: inference error: {e}", exc_info=True) - raise HTTPException(status_code=500, detail=f"Detection failed: {e}") + return { + "type": "classification", + "model": body.model, + "image_width": img_width, + "image_height": img_height, + "count": len(classifications), + "classifications": classifications, + } # ── Save Transformed Image ─────────────────────────────────────────────────── @router.post("/save/image", response_model=SaveImageToCollectionResponse) -@pxt_retry() def save_transformed_image(body: TransformRequest): """Apply transform at full resolution and save as a new image in Pixeltable.""" user_id = config.DEFAULT_USER_ID - try: - # Get original filename for naming the derivative - img_table = _get_pxt_table("image") - name_rows = ( - img_table.where((img_table.uuid == body.uuid) & (img_table.user_id == user_id)) - .select(img_source=img_table.image) - .collect() - ) - - original_name = "image" - if len(name_rows) > 0: - original_name = _source_to_filename(name_rows[0]["img_source"]) + # Get original filename for naming the derivative + img_table = _get_pxt_table("image") + name_rows = ( + img_table.where((img_table.uuid == body.uuid) & (img_table.user_id == user_id)) + .select(img_source=img_table.image) + .collect() + ) - result = _get_full_res_transform(body) + original_name = "image" + if len(name_rows) > 0: + original_name = _source_to_filename(name_rows[0]["img_source"]) - # Ensure the result is RGB/RGBA (save as PNG) - derived_name = _derive_filename(original_name, body.operation) - if not derived_name.lower().endswith(".png"): - derived_name = os.path.splitext(derived_name)[0] + ".png" + result = _get_full_res_transform(body) - # Write to the upload folder - os.makedirs(config.UPLOAD_FOLDER, exist_ok=True) - file_uuid = str(uuid_mod.uuid4()) - save_path = os.path.join(config.UPLOAD_FOLDER, f"{file_uuid}_{derived_name}") + # Ensure the result is RGB/RGBA (save as PNG) + derived_name = _derive_filename(original_name, body.operation) + if not derived_name.lower().endswith(".png"): + derived_name = os.path.splitext(derived_name)[0] + ".png" - if result.mode == "L": - result = result.convert("RGB") - result.save(save_path, format="PNG") + # Write to the upload folder + os.makedirs(config.UPLOAD_FOLDER, exist_ok=True) + file_uuid = str(uuid_mod.uuid4()) + save_path = os.path.join(config.UPLOAD_FOLDER, f"{file_uuid}_{derived_name}") - # Insert into Pixeltable images table - row = ImageRow( - image=save_path, - uuid=file_uuid, - timestamp=datetime.now(), - user_id=user_id, - ) - img_table.insert([row]) + if result.mode == "L": + result = result.convert("RGB") + result.save(save_path, format="PNG") - return SaveImageToCollectionResponse( - message=f"Saved {body.operation} result as new image", - uuid=file_uuid, - filename=derived_name, - ) + # Insert into Pixeltable images table + row = ImageRow( + image=save_path, + uuid=file_uuid, + timestamp=datetime.now(), + user_id=user_id, + ) + img_table.insert([row]) - except HTTPException: - raise - except Exception as e: - logger.error(f"Studio: save image error: {e}", exc_info=True) - raise HTTPException(status_code=500, detail=str(e)) + return SaveImageToCollectionResponse( + message=f"Saved {body.operation} result as new image", + uuid=file_uuid, + filename=derived_name, + ) # ── Download Transformed Image ─────────────────────────────────────────────── @@ -1753,44 +1660,37 @@ def save_transformed_image(body: TransformRequest): @pxt_retry() def download_transformed_image(body: TransformRequest): """Apply transform at full resolution and return as a downloadable PNG.""" - try: - # Get original filename for the download name - user_id = config.DEFAULT_USER_ID - img_table = _get_pxt_table("image") - name_rows = ( - img_table.where((img_table.uuid == body.uuid) & (img_table.user_id == user_id)) - .select(img_source=img_table.image) - .collect() - ) - - original_name = "image" - if len(name_rows) > 0: - original_name = _source_to_filename(name_rows[0]["img_source"]) + # Get original filename for the download name + user_id = config.DEFAULT_USER_ID + img_table = _get_pxt_table("image") + name_rows = ( + img_table.where((img_table.uuid == body.uuid) & (img_table.user_id == user_id)) + .select(img_source=img_table.image) + .collect() + ) - result = _get_full_res_transform(body) + original_name = "image" + if len(name_rows) > 0: + original_name = _source_to_filename(name_rows[0]["img_source"]) - derived_name = _derive_filename(original_name, body.operation) - if not derived_name.lower().endswith(".png"): - derived_name = os.path.splitext(derived_name)[0] + ".png" + result = _get_full_res_transform(body) - if result.mode == "L": - result = result.convert("RGB") + derived_name = _derive_filename(original_name, body.operation) + if not derived_name.lower().endswith(".png"): + derived_name = os.path.splitext(derived_name)[0] + ".png" - buf = io.BytesIO() - result.save(buf, format="PNG") - buf.seek(0) + if result.mode == "L": + result = result.convert("RGB") - return StreamingResponse( - buf, - media_type="image/png", - headers={"Content-Disposition": f'attachment; filename="{derived_name}"'}, - ) + buf = io.BytesIO() + result.save(buf, format="PNG") + buf.seek(0) - except HTTPException: - raise - except Exception as e: - logger.error(f"Studio: download image error: {e}", exc_info=True) - raise HTTPException(status_code=500, detail=str(e)) + return StreamingResponse( + buf, + media_type="image/png", + headers={"Content-Disposition": f'attachment; filename="{derived_name}"'}, + ) # ── Video Transform ────────────────────────────────────────────────────────── @@ -1801,198 +1701,189 @@ def download_transformed_image(body: TransformRequest): def transform_video(body: TransformRequest): """Apply a Pixeltable video UDF and return the result (metadata, frame, clip, overlay, scenes).""" user_id = config.DEFAULT_USER_ID - try: - vid_table = _get_pxt_table("video") - match = vid_table.where((vid_table.uuid == body.uuid) & (vid_table.user_id == user_id)) - - if body.operation == "view_metadata": - rows = match.select( - meta=pxt_video.get_metadata(vid_table.video), - dur=pxt_video.get_duration(vid_table.video), - ).collect() - if not rows: - raise HTTPException(status_code=404, detail="Video not found") - meta = rows[0].get("meta", {}) - dur = rows[0].get("dur") - streams = meta.get("streams", []) - video_stream: dict = next((s for s in streams if s.get("type") == "video"), {}) - return { - "operation": "view_metadata", - "duration": round(dur, 2) if dur else None, - "metadata": { - "format_size": meta.get("size"), - "bit_rate": meta.get("bit_rate"), - "width": video_stream.get("width"), - "height": video_stream.get("height"), - "fps": video_stream.get("average_rate"), - "total_frames": video_stream.get("frames"), - "codec": video_stream.get("codec_context", {}).get("name"), - "profile": video_stream.get("codec_context", {}).get("profile"), - "pix_fmt": video_stream.get("codec_context", {}).get("pix_fmt"), - }, - } - - elif body.operation == "extract_frame": - ts = float(body.params.get("timestamp", 0.0)) - rows = match.select( - frame=pxt_video.extract_frame(vid_table.video, timestamp=ts), - ).collect() - if not rows: - raise HTTPException(status_code=404, detail="Video not found") - frame = rows[0].get("frame") - if not isinstance(frame, Image.Image): - raise HTTPException(status_code=400, detail="No frame at that timestamp (may be past end of video)") - preview = _pil_image_to_data_uri(frame, max_size=PREVIEW_SIZE) - return { - "operation": "extract_frame", - "frame": preview, - "width": frame.size[0], - "height": frame.size[1], - "timestamp": ts, - } + vid_table = _get_pxt_table("video") + match = vid_table.where((vid_table.uuid == body.uuid) & (vid_table.user_id == user_id)) + + if body.operation == "view_metadata": + rows = match.select( + meta=pxt_video.get_metadata(vid_table.video), + dur=pxt_video.get_duration(vid_table.video), + ).collect() + if not rows: + raise HTTPException(status_code=404, detail="Video not found") + meta = rows[0].get("meta", {}) + dur = rows[0].get("dur") + streams = meta.get("streams", []) + video_stream: dict = next((s for s in streams if s.get("type") == "video"), {}) + return { + "operation": "view_metadata", + "duration": round(dur, 2) if dur else None, + "metadata": { + "format_size": meta.get("size"), + "bit_rate": meta.get("bit_rate"), + "width": video_stream.get("width"), + "height": video_stream.get("height"), + "fps": video_stream.get("average_rate"), + "total_frames": video_stream.get("frames"), + "codec": video_stream.get("codec_context", {}).get("name"), + "profile": video_stream.get("codec_context", {}).get("profile"), + "pix_fmt": video_stream.get("codec_context", {}).get("pix_fmt"), + }, + } - elif body.operation == "clip_video": - start = float(body.params.get("start", 0.0)) - duration = float(body.params.get("duration", 10.0)) - rows = match.select( - clipped=pxt_video.clip(vid_table.video, start_time=start, duration=duration), - ).collect() - if not rows: - raise HTTPException(status_code=404, detail="Video not found") - clipped = rows[0].get("clipped") - if clipped is None: - raise HTTPException(status_code=400, detail="Clip start is past end of video") - video_path = str(clipped) - clip_dur = round(duration, 2) - return { - "operation": "clip_video", - "video_url": f"/api/serve_video?path={video_path}", - "video_path": video_path, - "duration": clip_dur, - } + elif body.operation == "extract_frame": + ts = float(body.params.get("timestamp", 0.0)) + rows = match.select( + frame=pxt_video.extract_frame(vid_table.video, timestamp=ts), + ).collect() + if not rows: + raise HTTPException(status_code=404, detail="Video not found") + frame = rows[0].get("frame") + if not isinstance(frame, Image.Image): + raise HTTPException(status_code=400, detail="No frame at that timestamp (may be past end of video)") + preview = _pil_image_to_data_uri(frame, max_size=PREVIEW_SIZE) + return { + "operation": "extract_frame", + "frame": preview, + "width": frame.size[0], + "height": frame.size[1], + "timestamp": ts, + } - elif body.operation == "overlay_text": - text = str(body.params.get("text", "Hello World")) - font_size = int(body.params.get("font_size", 32)) - position = str(body.params.get("position", "bottom")) - v_align = "bottom" if position == "bottom" else "top" if position == "top" else "center" - rows = match.select( - result=pxt_video.overlay_text( - vid_table.video, - text, - font_size=font_size, - color="white", - vertical_align=v_align, - vertical_margin=40, - horizontal_align="center", - box=True, - box_color="black", - box_opacity=0.7, - box_border=[8, 16], - ), - ).collect() - if not rows: - raise HTTPException(status_code=404, detail="Video not found") - result = rows[0].get("result") - video_path = str(result) - return { - "operation": "overlay_text", - "video_url": f"/api/serve_video?path={video_path}", - "video_path": video_path, - } + elif body.operation == "clip_video": + start = float(body.params.get("start", 0.0)) + duration = float(body.params.get("duration", 10.0)) + rows = match.select( + clipped=pxt_video.clip(vid_table.video, start_time=start, duration=duration), + ).collect() + if not rows: + raise HTTPException(status_code=404, detail="Video not found") + clipped = rows[0].get("clipped") + if clipped is None: + raise HTTPException(status_code=400, detail="Clip start is past end of video") + video_path = str(clipped) + clip_dur = round(duration, 2) + return { + "operation": "clip_video", + "video_url": f"/api/serve_video?path={video_path}", + "video_path": video_path, + "duration": clip_dur, + } - elif body.operation == "crop_video": - x = int(body.params.get("x", 0)) - y = int(body.params.get("y", 0)) - w = int(body.params.get("width", 640)) - h = int(body.params.get("height", 480)) - bbox = [x, y, w, h] - rows = match.select( - cropped=pxt_video.crop(vid_table.video, bbox=bbox, bbox_format="xywh"), - ).collect() - if not rows: - raise HTTPException(status_code=404, detail="Video not found") - cropped = rows[0].get("cropped") - if cropped is None: - raise HTTPException(status_code=400, detail="Crop failed — check bbox against video dimensions") - video_path = str(cropped) - return { - "operation": "crop_video", - "video_url": f"/api/serve_video?path={video_path}", - "video_path": video_path, - "bbox": bbox, - } + elif body.operation == "overlay_text": + text = str(body.params.get("text", "Hello World")) + font_size = int(body.params.get("font_size", 32)) + position = str(body.params.get("position", "bottom")) + v_align = "bottom" if position == "bottom" else "top" if position == "top" else "center" + rows = match.select( + result=pxt_video.overlay_text( + vid_table.video, + text, + font_size=font_size, + color="white", + vertical_align=v_align, + vertical_margin=40, + horizontal_align="center", + box=True, + box_color="black", + box_opacity=0.7, + box_border=[8, 16], + ), + ).collect() + if not rows: + raise HTTPException(status_code=404, detail="Video not found") + result = rows[0].get("result") + video_path = str(result) + return { + "operation": "overlay_text", + "video_url": f"/api/serve_video?path={video_path}", + "video_path": video_path, + } - elif body.operation == "resize_video": - w = int(body.params.get("width", 640)) - h = int(body.params.get("height", 480)) - # pxt_video.resize uses 'size=(w, h)' per typical PIL/Pixeltable patterns or 'width=w, height=h' - # Based on Pixeltable PR #1210 and image.resize, we will use size=(w,h) or w,h. - # However some functions use kwargs `width=w, height=h`. Let's try size=[w, h] or let's use kwargs directly based on `image.resize` - rows = match.select( - resized=pxt_video.resize(vid_table.video, size=[w, h]), - ).collect() - if not rows: - raise HTTPException(status_code=404, detail="Video not found") - resized = rows[0].get("resized") - if resized is None: - raise HTTPException(status_code=400, detail="Resize failed") - video_path = str(resized) - return { - "operation": "resize_video", - "video_url": f"/api/serve_video?path={video_path}", - "video_path": video_path, - "dimensions": [w, h], - } + elif body.operation == "crop_video": + x = int(body.params.get("x", 0)) + y = int(body.params.get("y", 0)) + w = int(body.params.get("width", 640)) + h = int(body.params.get("height", 480)) + bbox = [x, y, w, h] + rows = match.select( + cropped=pxt_video.crop(vid_table.video, bbox=bbox, bbox_format="xywh"), + ).collect() + if not rows: + raise HTTPException(status_code=404, detail="Video not found") + cropped = rows[0].get("cropped") + if cropped is None: + raise HTTPException(status_code=400, detail="Crop failed — check bbox against video dimensions") + video_path = str(cropped) + return { + "operation": "crop_video", + "video_url": f"/api/serve_video?path={video_path}", + "video_path": video_path, + "bbox": bbox, + } - elif body.operation == "concat_videos": - uuid_str = body.params.get("uuids", "") - uuids_list = [u.strip() for u in uuid_str.split(",") if u.strip()] - if len(uuids_list) < 2: - raise HTTPException(status_code=400, detail="Provide at least 2 comma-separated UUIDs") - vid_table = _get_pxt_table("video") - concat_result = ( - vid_table.where(vid_table.uuid.isin(uuids_list)) - .select(concat=pxt_video.concat_videos_agg(vid_table.timestamp, vid_table.video)) - .collect() - ) - if not concat_result or concat_result[0].get("concat") is None: - raise HTTPException( - status_code=400, detail="Concat failed — ensure all videos share the same resolution" - ) - video_path = str(concat_result[0]["concat"]) - return { - "operation": "concat_videos", - "video_url": f"/api/serve_video?path={video_path}", - "video_path": video_path, - "source_uuids": uuids_list, - } + elif body.operation == "resize_video": + w = int(body.params.get("width", 640)) + h = int(body.params.get("height", 480)) + # pxt_video.resize uses 'size=(w, h)' per typical PIL/Pixeltable patterns or 'width=w, height=h' + # Based on Pixeltable PR #1210 and image.resize, we will use size=(w,h) or w,h. + # However some functions use kwargs `width=w, height=h`. Let's try size=[w, h] or let's use kwargs directly based on `image.resize` + rows = match.select( + resized=pxt_video.resize(vid_table.video, size=[w, h]), + ).collect() + if not rows: + raise HTTPException(status_code=404, detail="Video not found") + resized = rows[0].get("resized") + if resized is None: + raise HTTPException(status_code=400, detail="Resize failed") + video_path = str(resized) + return { + "operation": "resize_video", + "video_url": f"/api/serve_video?path={video_path}", + "video_path": video_path, + "dimensions": [w, h], + } - elif body.operation == "detect_scenes": - threshold = float(body.params.get("threshold", 27.0)) - rows = match.select( - scenes=pxt_video.scene_detect_content(vid_table.video, threshold=threshold), - dur=pxt_video.get_duration(vid_table.video), - ).collect() - if not rows: - raise HTTPException(status_code=404, detail="Video not found") - scenes = rows[0].get("scenes", []) - dur = rows[0].get("dur") - return { - "operation": "detect_scenes", - "scenes": scenes, - "total_duration": round(dur, 2) if dur else None, - "scene_count": len(scenes), - } + elif body.operation == "concat_videos": + uuid_str = body.params.get("uuids", "") + uuids_list = [u.strip() for u in uuid_str.split(",") if u.strip()] + if len(uuids_list) < 2: + raise HTTPException(status_code=400, detail="Provide at least 2 comma-separated UUIDs") + vid_table = _get_pxt_table("video") + concat_result = ( + vid_table.where(vid_table.uuid.isin(uuids_list)) + .select(concat=pxt_video.concat_videos_agg(vid_table.timestamp, vid_table.video)) + .collect() + ) + if not concat_result or concat_result[0].get("concat") is None: + raise HTTPException(status_code=400, detail="Concat failed — ensure all videos share the same resolution") + video_path = str(concat_result[0]["concat"]) + return { + "operation": "concat_videos", + "video_url": f"/api/serve_video?path={video_path}", + "video_path": video_path, + "source_uuids": uuids_list, + } - else: - raise HTTPException(status_code=400, detail=f"Unknown video operation: {body.operation}") + elif body.operation == "detect_scenes": + threshold = float(body.params.get("threshold", 27.0)) + rows = match.select( + scenes=pxt_video.scene_detect_content(vid_table.video, threshold=threshold), + dur=pxt_video.get_duration(vid_table.video), + ).collect() + if not rows: + raise HTTPException(status_code=404, detail="Video not found") + scenes = rows[0].get("scenes", []) + dur = rows[0].get("dur") + return { + "operation": "detect_scenes", + "scenes": scenes, + "total_duration": round(dur, 2) if dur else None, + "scene_count": len(scenes), + } - except HTTPException: - raise - except Exception as e: - logger.error(f"Studio: video transform error: {e}", exc_info=True) - raise HTTPException(status_code=500, detail=str(e)) + else: + raise HTTPException(status_code=400, detail=f"Unknown video operation: {body.operation}") # ── Save Video Transform Result ────────────────────────────────────────────── @@ -2005,74 +1896,58 @@ class SaveVideoRequest(BaseModel): @router.post("/save/video") -@pxt_retry() def save_video_result(body: SaveVideoRequest): """Save a video transform result (clip or overlay) as a new video in Pixeltable.""" user_id = config.DEFAULT_USER_ID - try: - # Re-run the transform to get the result video - result = transform_video(TransformRequest(uuid=body.uuid, operation=body.operation, params=body.params)) - video_path = result.get("video_path") - if not video_path or not os.path.exists(video_path): - raise HTTPException(status_code=400, detail="Operation did not produce a video file") + # Re-run the transform to get the result video + result = transform_video(TransformRequest(uuid=body.uuid, operation=body.operation, params=body.params)) + video_path = result.get("video_path") + if not video_path or not os.path.exists(video_path): + raise HTTPException(status_code=400, detail="Operation did not produce a video file") - import shutil + import shutil - os.makedirs(config.UPLOAD_FOLDER, exist_ok=True) - file_uuid = str(uuid_mod.uuid4()) - dest = os.path.join(config.UPLOAD_FOLDER, f"{file_uuid}_{body.operation}.mp4") - shutil.copy2(video_path, dest) + os.makedirs(config.UPLOAD_FOLDER, exist_ok=True) + file_uuid = str(uuid_mod.uuid4()) + dest = os.path.join(config.UPLOAD_FOLDER, f"{file_uuid}_{body.operation}.mp4") + shutil.copy2(video_path, dest) - vid_table = _get_pxt_table("video") - vid_table.insert([VideoRow(video=dest, uuid=file_uuid, timestamp=datetime.now(), user_id=user_id)]) + vid_table = _get_pxt_table("video") + vid_table.insert([VideoRow(video=dest, uuid=file_uuid, timestamp=datetime.now(), user_id=user_id)]) - return {"message": f"Saved {body.operation} result as new video", "uuid": file_uuid} - - except HTTPException: - raise - except Exception as e: - logger.error(f"Studio: save video error: {e}", exc_info=True) - raise HTTPException(status_code=500, detail=str(e)) + return {"message": f"Saved {body.operation} result as new video", "uuid": file_uuid} # ── Save Extracted Frame as Image ───────────────────────────────────────────── @router.post("/save/extracted_frame") -@pxt_retry() def save_extracted_frame(body: TransformRequest): """Extract a frame and save it as a new image in Pixeltable.""" user_id = config.DEFAULT_USER_ID - try: - vid_table = _get_pxt_table("video") - ts = float(body.params.get("timestamp", 0.0)) - rows = ( - vid_table.where((vid_table.uuid == body.uuid) & (vid_table.user_id == user_id)) - .select(frame=pxt_video.extract_frame(vid_table.video, timestamp=ts)) - .collect() - ) - - if not rows: - raise HTTPException(status_code=404, detail="Video not found") - frame = rows[0].get("frame") - if not isinstance(frame, Image.Image): - raise HTTPException(status_code=400, detail="No frame at that timestamp") + vid_table = _get_pxt_table("video") + ts = float(body.params.get("timestamp", 0.0)) + rows = ( + vid_table.where((vid_table.uuid == body.uuid) & (vid_table.user_id == user_id)) + .select(frame=pxt_video.extract_frame(vid_table.video, timestamp=ts)) + .collect() + ) - os.makedirs(config.UPLOAD_FOLDER, exist_ok=True) - file_uuid = str(uuid_mod.uuid4()) - save_path = os.path.join(config.UPLOAD_FOLDER, f"{file_uuid}_frame_{ts}s.png") - frame.save(save_path, format="PNG") + if not rows: + raise HTTPException(status_code=404, detail="Video not found") + frame = rows[0].get("frame") + if not isinstance(frame, Image.Image): + raise HTTPException(status_code=400, detail="No frame at that timestamp") - img_table = _get_pxt_table("image") - img_table.insert([ImageRow(image=save_path, uuid=file_uuid, timestamp=datetime.now(), user_id=user_id)]) + os.makedirs(config.UPLOAD_FOLDER, exist_ok=True) + file_uuid = str(uuid_mod.uuid4()) + save_path = os.path.join(config.UPLOAD_FOLDER, f"{file_uuid}_frame_{ts}s.png") + frame.save(save_path, format="PNG") - return {"message": f"Saved frame at {ts}s as new image", "uuid": file_uuid} + img_table = _get_pxt_table("image") + img_table.insert([ImageRow(image=save_path, uuid=file_uuid, timestamp=datetime.now(), user_id=user_id)]) - except HTTPException: - raise - except Exception as e: - logger.error(f"Studio: save extracted frame error: {e}", exc_info=True) - raise HTTPException(status_code=500, detail=str(e)) + return {"message": f"Saved frame at {ts}s as new image", "uuid": file_uuid} # ── Image Operations ───────────────────────────────────────────────────────── @@ -2156,36 +2031,27 @@ def _apply_image_operation(img: Image.Image, operation: str, params: dict) -> Im def get_document_summary(uuid: str): """Get the auto-generated summary for a document.""" user_id = config.DEFAULT_USER_ID - try: - doc_table = _get_pxt_table("document") - select_cols = dict(uuid_col=doc_table.uuid) - if hasattr(doc_table, "summary"): - select_cols["summary_json"] = doc_table.summary - if hasattr(doc_table, "document_text"): - select_cols["doc_text"] = doc_table.document_text + doc_table = _get_pxt_table("document") + select_cols = dict(uuid_col=doc_table.uuid) + if hasattr(doc_table, "summary"): + select_cols["summary_json"] = doc_table.summary + if hasattr(doc_table, "document_text"): + select_cols["doc_text"] = doc_table.document_text - rows = ( - doc_table.where((doc_table.uuid == uuid) & (doc_table.user_id == user_id)).select(**select_cols).collect() - ) - - if len(rows) == 0: - raise HTTPException(status_code=404, detail="Document not found") + rows = doc_table.where((doc_table.uuid == uuid) & (doc_table.user_id == user_id)).select(**select_cols).collect() - row = rows[0] - summary = _parse_summary(row.get("summary_json")) - doc_text_preview = (row.get("doc_text") or "")[:500] + if len(rows) == 0: + raise HTTPException(status_code=404, detail="Document not found") - return { - "uuid": uuid, - "summary": summary, - "text_preview": doc_text_preview, - } + row = rows[0] + summary = _parse_summary(row.get("summary_json")) + doc_text_preview = (row.get("doc_text") or "")[:500] - except HTTPException: - raise - except Exception as e: - logger.error(f"Studio: error fetching document summary: {e}", exc_info=True) - raise HTTPException(status_code=500, detail=str(e)) + return { + "uuid": uuid, + "summary": summary, + "text_preview": doc_text_preview, + } # ── Document Chunks ────────────────────────────────────────────────────────── @@ -2196,34 +2062,29 @@ def get_document_summary(uuid: str): def get_document_chunks(uuid: str, limit: int = 50): """Get extracted text chunks for a document.""" user_id = config.DEFAULT_USER_ID - try: - chunks_view = pxt.get_table("pixelbot_v3.chunks") - results = [] - for row in ( - chunks_view.where((chunks_view.uuid == uuid) & (chunks_view.user_id == user_id)) - .select( - text=chunks_view.text, - title=chunks_view.title, - heading=chunks_view.heading, - page=chunks_view.page, - ) - .limit(limit) - .collect() - ): - results.append( - { - "text": row.get("text", ""), - "title": row.get("title"), - "heading": row.get("heading"), - "page": row.get("page"), - } - ) - - return {"uuid": uuid, "chunks": results, "total": len(results)} + chunks_view = pxt.get_table("pixelbot_v3.chunks") + results = [] + for row in ( + chunks_view.where((chunks_view.uuid == uuid) & (chunks_view.user_id == user_id)) + .select( + text=chunks_view.text, + title=chunks_view.title, + heading=chunks_view.heading, + page=chunks_view.page, + ) + .limit(limit) + .collect() + ): + results.append( + { + "text": row.get("text", ""), + "title": row.get("title"), + "heading": row.get("heading"), + "page": row.get("page"), + } + ) - except Exception as e: - logger.error(f"Studio: error fetching chunks: {e}", exc_info=True) - raise HTTPException(status_code=500, detail=str(e)) + return {"uuid": uuid, "chunks": results, "total": len(results)} # ── Video Frames ───────────────────────────────────────────────────────────── @@ -2234,36 +2095,31 @@ def get_document_chunks(uuid: str, limit: int = 50): def get_video_frames(uuid: str, limit: int = 12): """Get extracted frames from a video as base64 thumbnails.""" user_id = config.DEFAULT_USER_ID - try: - frames_view = pxt.get_table("pixelbot_v3.video_frames") - results = [] - for row in ( - frames_view.where((frames_view.uuid == uuid) & (frames_view.user_id == user_id)) - .select( - frame=frames_view.frame, - pos_msec=frames_view.pos_msec, - ) - .order_by(frames_view.pos_msec) - .limit(limit) - .collect() - ): - frame = row.get("frame") - if isinstance(frame, Image.Image): - thumb = create_thumbnail_base64(frame, (192, 192)) - if thumb: - pos_sec = round(row.get("pos_msec", 0) / 1000, 1) - results.append( - { - "frame": thumb, - "position": pos_sec, - } - ) - - return {"uuid": uuid, "frames": results, "total": len(results)} + frames_view = pxt.get_table("pixelbot_v3.video_frames") + results = [] + for row in ( + frames_view.where((frames_view.uuid == uuid) & (frames_view.user_id == user_id)) + .select( + frame=frames_view.frame, + pos_msec=frames_view.pos_msec, + ) + .order_by(frames_view.pos_msec) + .limit(limit) + .collect() + ): + frame = row.get("frame") + if isinstance(frame, Image.Image): + thumb = create_thumbnail_base64(frame, (192, 192)) + if thumb: + pos_sec = round(row.get("pos_msec", 0) / 1000, 1) + results.append( + { + "frame": thumb, + "position": pos_sec, + } + ) - except Exception as e: - logger.error(f"Studio: error fetching frames: {e}", exc_info=True) - raise HTTPException(status_code=500, detail=str(e)) + return {"uuid": uuid, "frames": results, "total": len(results)} # ── Transcription ──────────────────────────────────────────────────────────── @@ -2277,27 +2133,21 @@ def get_transcription(uuid: str, media_type: str): if media_type not in ("audio", "video"): raise HTTPException(status_code=400, detail="media_type must be 'audio' or 'video'") - - try: - if media_type == "audio": - view_name = "pixelbot_v3.audio_chunks" - else: - view_name = "pixelbot_v3.video_audio_chunks" - - view = pxt.get_table(view_name) - sentences = [] - for row in view.where((view.uuid == uuid) & (view.user_id == user_id)).select(text=view.text).collect(): - text = row.get("text", "") - if text and text.strip(): - sentences.append(text.strip()) - - return { - "uuid": uuid, - "media_type": media_type, - "sentences": sentences, - "full_text": " ".join(sentences), - } - - except Exception as e: - logger.error(f"Studio: error fetching transcription: {e}", exc_info=True) - raise HTTPException(status_code=500, detail=str(e)) + if media_type == "audio": + view_name = "pixelbot_v3.audio_chunks" + else: + view_name = "pixelbot_v3.video_audio_chunks" + + view = pxt.get_table(view_name) + sentences = [] + for row in view.where((view.uuid == uuid) & (view.user_id == user_id)).select(text=view.text).collect(): + text = row.get("text", "") + if text and text.strip(): + sentences.append(text.strip()) + + return { + "uuid": uuid, + "media_type": media_type, + "sentences": sentences, + "full_text": " ".join(sentences), + } diff --git a/backend/pixelbot/utils.py b/backend/pixelbot/utils.py index 9008552..5ce24ea 100644 --- a/backend/pixelbot/utils.py +++ b/backend/pixelbot/utils.py @@ -1,9 +1,7 @@ # utils.py - Shared utility functions for the backend routers. -import asyncio import base64 import functools -import inspect import io import logging import os @@ -48,57 +46,34 @@ def pxt_retry( ) -> Callable[[Callable[P, T]], Callable[P, T]]: """Decorator that retries a function on transient Pixeltable connection errors. - Supports sync and async functions and retries only documented transient - database failures. Programming and assertion failures surface immediately. + Routes using this decorator are synchronous and read-only. Programming and + assertion failures surface immediately. """ def decorator(fn: Callable[P, T]) -> Callable[P, T]: - if inspect.iscoroutinefunction(fn): - - @functools.wraps(fn) - async def async_wrapper(*args: P.args, **kwargs: P.kwargs) -> T: - last_exc: Exception | None = None - wait = delay - for attempt in range(1, max_attempts + 1): - try: - return await fn(*args, **kwargs) - except Exception as exc: - if not _is_transient(exc) or attempt == max_attempts: - raise - last_exc = exc - logger.warning( - f"[pxt_retry] {fn.__name__} attempt {attempt}/{max_attempts} " - f"failed with transient error: {type(exc).__name__}: {str(exc)[:120]}. " - f"Retrying in {wait:.1f}s..." - ) - await asyncio.sleep(wait) - wait *= backoff - raise last_exc # type: ignore[misc] - - return async_wrapper # type: ignore[return-value] - else: - - @functools.wraps(fn) - def wrapper(*args: P.args, **kwargs: P.kwargs) -> T: - last_exc: Exception | None = None - wait = delay - for attempt in range(1, max_attempts + 1): - try: - return fn(*args, **kwargs) - except Exception as exc: - if not _is_transient(exc) or attempt == max_attempts: - raise - last_exc = exc - logger.warning( - f"[pxt_retry] {fn.__name__} attempt {attempt}/{max_attempts} " - f"failed with transient error: {type(exc).__name__}: {str(exc)[:120]}. " - f"Retrying in {wait:.1f}s..." - ) - time.sleep(wait) - wait *= backoff - raise last_exc # type: ignore[misc] - - return wrapper + @functools.wraps(fn) + def wrapper(*args: P.args, **kwargs: P.kwargs) -> T: + wait = delay + for attempt in range(1, max_attempts + 1): + try: + return fn(*args, **kwargs) + except Exception as exc: + if not _is_transient(exc) or attempt == max_attempts: + raise + logger.warning( + "[pxt_retry] %s attempt %d/%d failed with transient %s: %s. Retrying in %.1fs...", + fn.__name__, + attempt, + max_attempts, + type(exc).__name__, + str(exc)[:120], + wait, + ) + time.sleep(wait) + wait *= backoff + raise RuntimeError("unreachable") + + return wrapper return decorator diff --git a/backend/pyproject.toml b/backend/pyproject.toml index 8512fac..9783113 100644 --- a/backend/pyproject.toml +++ b/backend/pyproject.toml @@ -14,6 +14,7 @@ dependencies = [ "anthropic", "openai", "mistralai", + # Required by the core CLIP image/video indexes as well as Studio detection. "torch", "transformers", "spacy", diff --git a/backend/tests/conftest.py b/backend/tests/conftest.py index dcab5be..f68e16e 100644 --- a/backend/tests/conftest.py +++ b/backend/tests/conftest.py @@ -1,5 +1,26 @@ import os +import subprocess import tempfile +from importlib.util import find_spec +from pathlib import Path # Pixeltable reads this while modules are imported during test collection. -os.environ["PIXELTABLE_HOME"] = tempfile.mkdtemp(prefix="pixelbot-test-") +TEST_CATALOG = Path(tempfile.mkdtemp(prefix="pixelbot-test-")) +os.environ["PIXELTABLE_HOME"] = str(TEST_CATALOG) + + +def pytest_sessionfinish(session, exitstatus) -> None: + """Stop the temporary PostgreSQL server that Pixeltable keeps persistent.""" + pgdata = TEST_CATALOG / "pgdata" + if not (pgdata / "postmaster.pid").exists(): + return + package = find_spec("pixeltable_pgserver") + if package is None or package.submodule_search_locations is None: + return + executable = Path(next(iter(package.submodule_search_locations))) / "pginstall" / "bin" / "pg_ctl" + subprocess.run( + [str(executable), "-D", str(pgdata), "stop", "-m", "fast"], + capture_output=True, + check=False, + timeout=30, + ) diff --git a/backend/tests/test_contract.py b/backend/tests/test_contract.py index 5472064..6a5149e 100644 --- a/backend/tests/test_contract.py +++ b/backend/tests/test_contract.py @@ -2,12 +2,13 @@ import io import socket from pathlib import Path +from types import SimpleNamespace import pytest from fastapi import HTTPException, UploadFile from fastapi.testclient import TestClient -from pixelbot import __version__, config +from pixelbot import __version__, config, notifications from pixelbot.app import app from pixelbot.catalog_access import require_allowed_table from pixelbot.functions import send_webhook @@ -37,6 +38,13 @@ def test_removed_routes_are_not_registered() -> None: "/api/memory/v2/delete", "/api/memory/manual", "/api/download_memory", + "/api/delete_file/{file_uuid}/{file_type}", + "/api/delete_all", + "/api/workflow_detail/{timestamp_str}", + "/api/delete_history/{timestamp_str}", + "/api/tts_voices", + "/api/db/table/{path}/schema", + "/api/db/table/{path}/versions", } assert not any(path in removed for _, path in routes) @@ -55,14 +63,50 @@ def test_pixeltable_query_routes_are_canonical() -> None: def test_removed_api_routes_return_not_found() -> None: client = TestClient(app) - assert client.post("/api/studio/reve/edit", json={}).status_code == 404 - assert client.post("/api/db/create_table", json={}).status_code == 404 + removed_requests = ( + ("POST", "/api/studio/reve/edit"), + ("POST", "/api/db/create_table"), + ("DELETE", "/api/delete_file/example/image"), + ("POST", "/api/delete_all"), + ("GET", "/api/workflow_detail/2026-01-01"), + ("DELETE", "/api/delete_history/2026-01-01"), + ("GET", "/api/tts_voices"), + ("GET", "/api/db/table/pixelbot_v3.images/schema"), + ("GET", "/api/db/table/pixelbot_v3.images/versions"), + ) + for method, path in removed_requests: + assert client.request(method, path).status_code == 404 + + +def test_unexpected_route_errors_are_sanitized(monkeypatch: pytest.MonkeyPatch) -> None: + def fail_catalog_lookup(_: str): + raise RuntimeError("secret catalog detail") + + monkeypatch.setattr("pixelbot.routers.history.pxt.get_table", fail_catalog_lookup) + response = TestClient(app, raise_server_exceptions=False).get("/api/conversations") + + assert response.status_code == 500 + assert response.json() == {"detail": "Internal server error"} def test_webhook_destination_is_configuration_only() -> None: assert list(inspect.signature(send_webhook.py_fn).parameters) == ["message"] +def test_notification_delivery_is_shared_and_typed(monkeypatch: pytest.MonkeyPatch) -> None: + monkeypatch.setattr(config, "SLACK_WEBHOOK_URL", "https://example.test/secret-token") + monkeypatch.setattr( + notifications.requests, + "post", + lambda url, **kwargs: SimpleNamespace(status_code=200), + ) + + result = notifications.deliver_notification("slack", "hello") + + assert result == notifications.DeliveryResult("Slack message sent successfully.", True, 200) + assert notifications.redacted_destination("slack") == "https://example.test/..." + + def test_private_url_is_rejected(monkeypatch: pytest.MonkeyPatch) -> None: monkeypatch.setattr( socket, "getaddrinfo", lambda *_: [(socket.AF_INET, socket.SOCK_STREAM, 6, "", ("127.0.0.1", 80))] diff --git a/backend/tests/test_schema_cli.py b/backend/tests/test_schema_cli.py index e6540f7..5fafe47 100644 --- a/backend/tests/test_schema_cli.py +++ b/backend/tests/test_schema_cli.py @@ -5,9 +5,8 @@ from pathlib import Path -def test_schema_check_uses_an_isolated_catalog(tmp_path: Path) -> None: +def test_schema_check_uses_an_isolated_catalog() -> None: env = os.environ.copy() - env["PIXELTABLE_HOME"] = str(tmp_path / "catalog") env["PYTHONPATH"] = str(Path(__file__).parents[1]) with socket.socket() as port_socket: port_socket.bind(("127.0.0.1", 0)) diff --git a/docs/pixeltable-0.7.7-upgrade.md b/docs/pixeltable-0.7.7-upgrade.md index 8042162..49d3b9a 100644 --- a/docs/pixeltable-0.7.7-upgrade.md +++ b/docs/pixeltable-0.7.7-upgrade.md @@ -57,4 +57,10 @@ A fresh local schema apply completed with 19 models. A second diff reported all On September 12, 2026, the read API was simplified against the canonical Pixeltable app guidance. The duplicated `pixelbot/queries.py` execution layer and versioned memory/persona reads were removed. `pixelbot/app.py` now declares the endpoint queries and exposes them through `FastAPIRouter`; custom FastAPI handlers retain the validation-heavy writes. A catalog created from the merged 3.0 schema remained fully in agreement (19 of 19 models, zero schema operations), and a real managed service returned 200 from the canonical memory and persona reads while the removed versioned routes returned 404. +A second simplification removed unused file-wide deletion, per-entry workflow, TTS voice-list, and standalone database schema/version endpoints. Duplicate route-level exception wrappers now fall through to the application's sanitized error handler, while expected validation and missing-resource responses remain explicit. Retries remain only on read-only operations so provider calls and catalog writes cannot be repeated after partial success. Notification delivery is implemented once for both HTTP tests and Pixeltable tools, returns typed delivery status, and stores a redacted destination origin. + +The pytest fixture now reuses one isolated Pixeltable catalog and stops its PostgreSQL server when the session ends. Two consecutive full backend runs passed with 15 tests and left no Pixelbot test database process running. + +Torch and Transformers remain base dependencies because the declared CLIP indexes power core image and video-frame retrieval. Making Studio detection optional would not reduce the installation until those indexes move to another embedding model; that change requires retrieval-quality evaluation and similarity-threshold retuning and is outside this behavior-preserving cleanup. + Provider calls are mocked and establish wiring only. Paid provider calls, hosted deployment, Cloud behavior, and the 48 agent trials were not run. Validation used a separate temporary `PIXELTABLE_HOME`; the old `agents` catalog was not changed. diff --git a/frontend/src/types/index.ts b/frontend/src/types/index.ts index 6c0436d..25cc546 100644 --- a/frontend/src/types/index.ts +++ b/frontend/src/types/index.ts @@ -96,12 +96,6 @@ export interface Conversation { message_count: number } -export interface TtsVoice { - id: string - label: string - style: string -} - export interface JoinResult { left_table: string right_table: string @@ -527,24 +521,6 @@ export interface PipelineResponse { edges: PipelineEdge[] } -export interface VersionEntry { - version: number - created_at: string | null - change_type: string | null - inserts: number - updates: number - deletes: number - errors: number - schema_change?: string | null -} - -export interface VersionsResponse { - path: string - current_version: number - can_revert: boolean - versions: VersionEntry[] -} - // Integrations export interface IntegrationInfo { id: string