Spaces:
Running
Running
| from __future__ import annotations | |
| import json | |
| from typing import Annotated, Any, Dict, List, Optional | |
| from fastapi import APIRouter, File, Form, HTTPException, Request, UploadFile | |
| from pydantic import BaseModel, ValidationError | |
| from app.api.deps import get_redis_scripts | |
| from app.config import get_settings | |
| from app.services.chat_service import chat_completion | |
| from app.services.csv_analysis_service import ( | |
| execute_csv_chat_blocks, | |
| get_dataset_info, | |
| ) | |
| from app.services.prompts import get_csv_system_prompt | |
| from app.utils.json_utils import extract_json_blocks | |
| class _AnalyzeBlock(BaseModel): | |
| description: str = "" | |
| python_code: str = "" | |
| class _VisualizationBlock(BaseModel): | |
| description: str = "" | |
| python_code: str = "" | |
| class _AIResponse(BaseModel): | |
| analyze: List[_AnalyzeBlock] = [] | |
| visualization: List[_VisualizationBlock] = [] | |
| message: str = "" | |
| router = APIRouter() | |
| _settings = get_settings() | |
| _MAX_UPLOAD_BYTES = _settings.max_upload_bytes | |
| async def get_csv_info( | |
| files: Annotated[Optional[List[UploadFile]], File(description="CSV files to inspect (max 10 total with URLs)")] = None, | |
| urls: Annotated[Optional[str], Form(description="JSON array of file URLs (max 10 total with files)")] = None, | |
| ): | |
| parsed_urls: List[str] = [] | |
| if urls: | |
| try: | |
| parsed_urls = json.loads(urls) | |
| if not isinstance(parsed_urls, list) or not all(isinstance(u, str) for u in parsed_urls): | |
| raise ValueError("urls must be a JSON array of strings") | |
| except (json.JSONDecodeError, ValueError) as exc: | |
| raise HTTPException(status_code=400, detail=str(exc)) | |
| file_count = len(files) if files else 0 | |
| url_count = len(parsed_urls) | |
| total = file_count + url_count | |
| if total == 0: | |
| raise HTTPException(status_code=400, detail="Provide at least one file or URL") | |
| if total > 10: | |
| raise HTTPException(status_code=400, detail=f"Maximum 10 sources allowed (got {total})") | |
| results: List[dict] = [] | |
| if files: | |
| for f in files: | |
| try: | |
| data = await f.read() | |
| except Exception as exc: | |
| results.append({"source": getattr(f, "filename", "unknown"), "success": False, "error": f"Read error: {exc}"}) | |
| continue | |
| if len(data) > _MAX_UPLOAD_BYTES: | |
| results.append({"source": f.filename or "unknown", "success": False, "error": f"File exceeds {_settings.max_upload_mb} MB limit"}) | |
| continue | |
| if not data: | |
| results.append({"source": f.filename or "unknown", "success": False, "error": "Empty file"}) | |
| continue | |
| try: | |
| meta = await get_dataset_info(data) | |
| meta["source"] = f.filename or "upload" | |
| results.append(meta) | |
| except Exception as exc: | |
| results.append({"source": f.filename or "upload", "success": False, "error": str(exc)}) | |
| for url in parsed_urls: | |
| if not url.startswith(("http://", "https://")): | |
| results.append({"source": url, "success": False, "error": "Only http/https URLs are supported"}) | |
| continue | |
| try: | |
| meta = await get_dataset_info(url) | |
| meta["source"] = url | |
| results.append(meta) | |
| except Exception as exc: | |
| results.append({"source": url, "success": False, "error": str(exc)}) | |
| return { | |
| "success": True, | |
| "total": total, | |
| "succeeded": sum(1 for r in results if r.get("success")), | |
| "failed": sum(1 for r in results if not r.get("success")), | |
| "results": results, | |
| } | |
| # @router.post( | |
| # "/csv/analyze", | |
| # summary="Execute Python analysis code against a CSV file (upload or URL)", | |
| # ) | |
| # async def analyze_csv( | |
| # file: Annotated[Optional[UploadFile], File(description="CSV file to analyze")] = None, | |
| # url: Annotated[Optional[str], Form(description="URL to a CSV file")] = None, | |
| # code: str = Form(..., description="Python code to execute (df pre-loaded with CSV data)"), | |
| # token: str = Depends(require_auth), | |
| # ): | |
| # if not file and not url: | |
| # raise HTTPException(status_code=400, detail="Provide either a file or a URL") | |
| # if file and url: | |
| # raise HTTPException(status_code=400, detail="Provide either a file or a URL, not both") | |
| # if file: | |
| # data = await file.read() | |
| # if len(data) > _MAX_UPLOAD_BYTES: | |
| # raise HTTPException(status_code=413, detail=f"File exceeds {_settings.max_upload_mb} MB limit") | |
| # if not data: | |
| # raise HTTPException(status_code=400, detail="Empty file") | |
| # result = await analyze_csv_dataset(data, code) | |
| # else: | |
| # result = await analyze_csv_dataset(url, code) | |
| # return result | |
| # @router.post( | |
| # "/csv/chart", | |
| # summary="Generate a chart from CSV data and return as base64 PNG (upload or URL)", | |
| # ) | |
| # async def chart_csv( | |
| # file: Annotated[Optional[UploadFile], File(description="CSV file for chart generation")] = None, | |
| # url: Annotated[Optional[str], Form(description="URL to a CSV file")] = None, | |
| # code: str = Form(..., description="Python chart code (df pre-loaded, use matplotlib/seaborn)"), | |
| # token: str = Depends(require_auth), | |
| # ): | |
| # if not file and not url: | |
| # raise HTTPException(status_code=400, detail="Provide either a file or a URL") | |
| # if file and url: | |
| # raise HTTPException(status_code=400, detail="Provide either a file or a URL, not both") | |
| # if file: | |
| # data = await file.read() | |
| # if len(data) > _MAX_UPLOAD_BYTES: | |
| # raise HTTPException(status_code=413, detail=f"File exceeds {_settings.max_upload_mb} MB limit") | |
| # if not data: | |
| # raise HTTPException(status_code=400, detail="Empty file") | |
| # result = await create_csv_chart(data, code) | |
| # else: | |
| # result = await create_csv_chart(url, code) | |
| # return result | |
| async def csv_chat( | |
| request: Request, | |
| file: Annotated[Optional[UploadFile], File(description="CSV file to analyze")] = None, | |
| ): | |
| content_type = request.headers.get("content-type", "") | |
| url: Optional[str] = None | |
| query: Optional[str] = None | |
| if "application/json" in content_type: | |
| try: | |
| body = await request.json() | |
| url = body.get("url") | |
| query = body.get("query") | |
| except Exception as exc: | |
| raise HTTPException(status_code=400, detail=f"Invalid JSON body: {exc}") | |
| else: | |
| form = await request.form() | |
| url = form.get("url") | |
| query = form.get("query") | |
| has_file = file is not None | |
| has_url = bool(url) | |
| if not has_file and not has_url: | |
| raise HTTPException(status_code=400, detail="Provide either a file or a URL") | |
| if has_file and has_url: | |
| raise HTTPException(status_code=400, detail="Provide either a file or a URL, not both") | |
| if has_file: | |
| data = await file.read() | |
| if len(data) > _MAX_UPLOAD_BYTES: | |
| raise HTTPException(status_code=413, detail=f"File exceeds {_settings.max_upload_mb} MB limit") | |
| if not data: | |
| raise HTTPException(status_code=400, detail="Empty file") | |
| source: Any = data | |
| else: | |
| source = url | |
| if not query: | |
| raise HTTPException(status_code=400, detail="Query is required") | |
| metadata = await get_dataset_info(source) | |
| system_prompt = get_csv_system_prompt(metadata) | |
| messages = [ | |
| {"role": "system", "content": system_prompt}, | |
| {"role": "user", "content": query}, | |
| ] | |
| redis, scripts = get_redis_scripts(request) | |
| try: | |
| ai_response = await chat_completion( | |
| messages=messages, | |
| response_format={"type": "json_object"}, | |
| max_tokens=12000, | |
| redis=redis, | |
| scripts=scripts, | |
| ) | |
| except RuntimeError as e: | |
| raise HTTPException(status_code=502, detail=str(e)) | |
| parsed = ai_response.get("parsed") | |
| if not parsed: | |
| choices = ai_response.get("choices", []) | |
| content = choices[0].get("message", {}).get("content", "") if choices else "" | |
| blocks = extract_json_blocks(content) | |
| if blocks: | |
| parsed = blocks[0] | |
| else: | |
| try: | |
| parsed = json.loads(content) | |
| except (json.JSONDecodeError, TypeError): | |
| pass | |
| if not isinstance(parsed, dict): | |
| return { | |
| "success": False, | |
| "message": None, | |
| "analyze": [], | |
| "visualizations": [], | |
| "error": "AI response was not valid JSON", | |
| } | |
| try: | |
| ai_data = _AIResponse(**parsed) | |
| except ValidationError as exc: | |
| return { | |
| "success": False, | |
| "message": None, | |
| "analyze": [], | |
| "visualizations": [], | |
| "error": f"AI response failed schema validation: {exc}", | |
| } | |
| message_text = ai_data.message | |
| has_content = bool(message_text.strip()) if message_text else False | |
| analyze_blocks_raw = [b.model_dump() for b in ai_data.analyze] | |
| viz_blocks_raw = [b.model_dump() for b in ai_data.visualization] | |
| exec_result = await execute_csv_chat_blocks( | |
| source=source, | |
| analyze_blocks=analyze_blocks_raw, | |
| viz_blocks=viz_blocks_raw, | |
| ) | |
| if not exec_result["success"]: | |
| return { | |
| "success": False, | |
| "message": ai_data.message if has_content else None, | |
| "analyze": [], | |
| "visualizations": [], | |
| "error": exec_result.get("error", "Code execution failed"), | |
| } | |
| results = exec_result.get("results", {}) | |
| raw_analyze = results.get("analyze", []) | |
| raw_visualizations = results.get("visualization", []) | |
| analyze_results: List[Dict[str, Any]] = [] | |
| for i, block in enumerate(ai_data.analyze): | |
| raw = raw_analyze[i] if i < len(raw_analyze) else {} | |
| code = block.python_code.strip() | |
| if not code: | |
| continue | |
| analyze_results.append({ | |
| "description": block.description, | |
| "code": code, | |
| "success": raw.get("success", False), | |
| "output": raw.get("output", ""), | |
| "error": raw.get("error"), | |
| "execution_time_ms": exec_result["execution_time_ms"], | |
| }) | |
| viz_results: List[Dict[str, Any]] = [] | |
| for i, block in enumerate(ai_data.visualization): | |
| raw = raw_visualizations[i] if i < len(raw_visualizations) else {} | |
| code = block.python_code.strip() | |
| if not code: | |
| continue | |
| viz_results.append({ | |
| "description": block.description, | |
| "code": code, | |
| "success": raw.get("success", False), | |
| "image_base64": raw.get("image_base64"), | |
| "error": raw.get("error"), | |
| "execution_time_ms": exec_result["execution_time_ms"], | |
| }) | |
| return { | |
| "success": True, | |
| "message": ai_data.message if has_content else None, | |
| "analyze": analyze_results, | |
| "visualizations": viz_results, | |
| "error": None, | |
| } | |