diff --git a/.env.example b/.env.example new file mode 100644 index 0000000..fb36f80 --- /dev/null +++ b/.env.example @@ -0,0 +1,60 @@ +# PremSQL Environment Configuration +# Copy this file to .env and fill in your values + +# ============================================================ +# Security Configuration +# ============================================================ + +# MODE SELECTION: +# - Development: PREMSQL_DJANGO_DEBUG=true (no token required) +# - Production: PREMSQL_DJANGO_DEBUG=false (token REQUIRED) + +# For development (recommended for local testing): +PREMSQL_DJANGO_DEBUG=true + +# For production, uncomment and set these: +#PREMSQL_DJANGO_DEBUG=false +#PREMSQL_API_TOKEN= +#PREMSQL_DJANGO_SECRET_KEY= + +# Optional: Custom session name (auto-generated if not set) +#PREMSQL_SESSION_NAME=my-session + +# ============================================================ +# Django Settings +# ============================================================ +#PREMSQL_DJANGO_DEBUG=false +#PREMSQL_ALLOWED_HOSTS=127.0.0.1,localhost +#PREMSQL_ALLOWED_ORIGINS=http://127.0.0.1:8501,http://localhost:8501 + +# ============================================================ +# LLM Provider Configuration +# Choose ONE provider from the options below +# ============================================================ + +# Option 1: vLLM (self-hosted models with OpenAI-compatible API) +# Start vLLM: vllm serve /path/to/model --port 8000 +#VLLM_BASE_URL=http://localhost:8000/v1 +#VLLM_MODEL_NAME=your-model-name + +# Option 2: Custom OpenAI-Compatible Service +# For LM Studio, LocalAI, Text Generation WebUI, or any custom deployment +#CUSTOM_BASE_URL=http://localhost:1234/v1 +#CUSTOM_MODEL_NAME=your-model-name +#CUSTOM_API_KEY=your-key-if-required + +# Option 3: OpenAI (official API) +# Get your API key from https://platform.openai.com/ +#OPENAI_API_KEY=sk-your-openai-api-key +#OPENAI_MODEL_NAME=gpt-4o-mini + +# Option 4: PremAI +# Get your API key from https://app.premai.io/ +#PREMAI_API_KEY=your-premai-api-key +#PREMAI_PROJECT_ID=your-project-id + +# Option 5: Ollama (local models) +# Install Ollama: https://ollama.com +# Run: ollama pull llama3.2 +#OLLAMA_BASE_URL=http://127.0.0.1:11434 +#OLLAMA_MODEL_NAME=llama3.2 \ No newline at end of file diff --git a/PR_DESCRIPTION.md b/PR_DESCRIPTION.md new file mode 100644 index 0000000..1d04048 --- /dev/null +++ b/PR_DESCRIPTION.md @@ -0,0 +1,201 @@ +## Summary + +This PR addresses security vulnerabilities (#40), fixes dependency conflicts (#37), and adds new features for better model support and deployment experience. + +--- + +## Dependency Fixes (Issue #37) + +Fixed version conflicts reported during installation: + +| Conflict | Cause | Resolution | +| ------------------------------------------------------- | ---------------------------------- | ------------------------ | +| `httpx<0.29` vs `>=0.27` required by ollama/browser-use | fastapi 0.112 pinned old httpx | Update fastapi >=0.115.0 | +| `starlette>=0.41.3` vs 0.38.6 | fastapi 0.112 pinned old starlette | Update fastapi >=0.115.0 | +| `python==3.12 not compatible` | Version range `^3.10` | Extend to `>=3.10,<3.13` | + +**Changes in pyproject.toml:** + +```toml +# Before +python = "^3.10" +fastapi = "^0.112.0" + +# After +python = ">=3.10,<3.13" # Support 3.11, 3.12 +fastapi = ">=0.115.0" # Brings httpx>=0.27, starlette>=0.41 +httpx = ">=0.27.0" # Explicit for clarity +starlette = ">=0.41.0" # Explicit for clarity +``` + +--- + +## Security Fixes (Issue #40) + +### Critical Severity +- **RCE via eval()** → Replace with `ast.literal_eval()` +- **SSRF** → `normalize_base_url()` restricts to loopback addresses +- **Missing Authentication** → Token-based auth via `PREMSQL_API_TOKEN` + +### High Severity +- **SQL Write Operations** → `enforce_read_only_sql()` blocks INSERT/UPDATE/DELETE/DROP +- **Path Traversal** → Whitelist validation: `[A-Za-z0-9_-]{1,64}` +- **Information Disclosure** → Removed sensitive fields from API responses +- **Error Message Leakage** → `safe_error_message()` whitelist mechanism + +### Medium Severity +- **Pickle RCE** → Add `weights_only=True` to `torch.load()` +- **Swagger Exposure** → Restrict access based on DEBUG mode +- **Process Management** → PID file-based precise process tracking +- **SQL Injection** → Parameterized queries with `?` placeholders + +--- + +## New Features + +### 1. LLM Provider Support + +Added new generators for self-hosted and custom LLM deployments: + +| Generator | Use Case | +| ----------------------------------- | -------------------------------------------------------- | +| `Text2SQLGeneratorVLLM` | vLLM deployed models (auto Qwen3 thinking mode handling) | +| `Text2SQLGeneratorOpenAICompatible` | Any OpenAI-compatible API (LM Studio, LocalAI, etc.) | + +**Usage:** +```python +from premsql.generators import Text2SQLGeneratorVLLM + +generator = Text2SQLGeneratorVLLM( + model_name="/models/qwen", + base_url="http://localhost:8000/v1", + experiment_name="test", type="test" +) +``` + +### 2. Easy Deployment Script + +Added `start_agent.py` for one-command AgentServer startup: + +```bash +# Configure in .env, then: +python start_agent.py +``` + +Auto-detects configured LLM provider from environment variables. + +### 3. Bug Fixes + +| Fix | Description | +| ----------------------- | ------------------------------------------------------------------------------ | +| Plot image generation | Fixed `plot_image=False` hardcoded in server_mode, now generates base64 images | +| Matplotlib backend | Added `Agg` backend for non-interactive server mode | +| Session duplicate error | Auto-delete existing session when creating new one with same name | +| Session deletion memory | Recreate memory DB table after deletion to prevent OperationalError | +| API token propagation | Fixed Django backend not passing API token to AgentServer | +| Environment loading | Added dotenv loading in Django manage.py and Streamlit main.py | + +### 4. UI Improvements + +- Session list now shows delete button for each session (no need to type session name) +- Delete operation auto-refreshes the page +- Simplified session creation form with placeholder hints + +--- + +## Configuration + +### Environment Variables (.env) + +All configuration is optional for local development - tokens are auto-generated if not set. + +```bash +# Security (auto-generated for local dev) +#PREMSQL_API_TOKEN=your-token-here +#PREMSQL_DJANGO_SECRET_KEY=your-secret-here + +# LLM Provider (choose one) +VLLM_BASE_URL=http://localhost:8000/v1 +VLLM_MODEL_NAME=/models/your-model + +# Or custom OpenAI-compatible service +CUSTOM_BASE_URL=http://localhost:1234/v1 +CUSTOM_MODEL_NAME=local-model + +# Or official OpenAI +OPENAI_API_KEY=sk-your-key +OPENAI_MODEL_NAME=gpt-4o-mini +``` + +### Quick Start + +```bash +# 1. Copy and configure .env +cp .env.example .env +# Edit .env to set your LLM provider + +# 2. Start services +python start_agent.py # AgentServer on port 8100 +premsql launch all # Django + Streamlit + +# 3. Open browser +http://localhost:8501 # PremSQL Playground +``` + +--- + +## Security Configuration + +### Development Mode (Recommended for Local Testing) + +Set `PREMSQL_DJANGO_DEBUG=true` to enable development mode: + +```bash +# .env for development +PREMSQL_DJANGO_DEBUG=true +# No need to set PREMSQL_API_TOKEN - auto-generated +``` + +In development mode: +- API token is auto-generated (printed in startup logs) +- Authentication is skipped for all services +- Suitable for local testing only + +### Production Mode (Required for Deployment) + +**IMPORTANT**: Production mode requires explicit token configuration: + +```bash +# .env for production +PREMSQL_DJANGO_DEBUG=false +PREMSQL_API_TOKEN= # REQUIRED! +PREMSQL_DJANGO_SECRET_KEY= +``` + +**How to generate secure tokens:** + +```bash +# Using Python (recommended) +python -c "import secrets; print(secrets.token_hex(32))" + +# Using OpenSSL +openssl rand -hex 32 + +# Using UUID +python -c "import uuid; print(uuid.uuid4().hex)" +``` + +Example output: `a1b2c3d4e5f6...` (64 characters hex string) + +**Security behavior summary:** + +| Mode | PREMSQL_API_TOKEN | Behavior | +| ----------- | ----------------- | ---------------------------- | +| Development | Not set | Auto-generated, auth skipped | +| Development | Set | Use configured token | +| Production | Not set | **Service fails to start** | +| Production | Set | Enforce authentication | + +--- + +Resolves #37 #40 diff --git a/README.md b/README.md index 650ab17..196268a 100644 --- a/README.md +++ b/README.md @@ -131,7 +131,11 @@ plot = agent( You can launch the PremSQL Playground (as shown in the above video by adding these two additional lines after instantiating Agent) ```python -agent_server = AgentServer(agent=agent, port={port}) +agent_server = AgentServer( + agent=agent, + port={port}, + api_token=os.environ.get("PREMSQL_API_TOKEN") +) agent_server.launch() ``` @@ -147,6 +151,24 @@ and on the second side of the terminal write: python start_agent.py ``` +### Security Defaults + +Recent hardening changes added safer defaults for Playground and AgentServer deployments: + +- Playground session registration only accepts loopback AgentServer URLs such as `http://127.0.0.1:8100`. +- API authentication is supported through `PREMSQL_API_TOKEN`. When this variable is set, the Django backend, Playground clients, and FastAPI AgentServer all expect the same token. +- Generated SQL is now restricted to read-only queries. `INSERT`, `UPDATE`, `DELETE`, `DROP`, and similar statements are rejected by default. +- Direct raw SQL passthrough using backtick-wrapped prompts is disabled by default. +- CSV imports now enforce upload limits to reduce accidental or malicious resource exhaustion. + +Recommended local environment variables: + +```bash +export PREMSQL_API_TOKEN="change-this-token" +export PREMSQL_DJANGO_SECRET_KEY="change-this-secret" +export PREMSQL_ALLOWED_HOSTS="127.0.0.1,localhost" +``` + ## 📦 Components Overview ### [Datasets](https://docs.premai.io/premsql/introduction) @@ -447,13 +469,18 @@ In the above section you have see how we have defined our agent. You can deploy ```python # File name: start_agent_server.py +import os from premsql.playground import AgentServer from premsql.agents import BaseLineAgent # Define your agent as shown above: agent = BaseLineAgent(...) -agent_server = AgentServer(agent=agent, port={port}) +agent_server = AgentServer( + agent=agent, + port={port}, + api_token=os.environ.get("PREMSQL_API_TOKEN") +) agent_server.launch() ``` @@ -463,7 +490,7 @@ Now inside another terminal write: python start_agent_server.py ``` -This can be any python file name. This will run a fastapi server. You need to paste the deployed url and paste it inside `Register New Session` part of the UI. Below shows, how the basic backend architecture looks like on how Playground communicates with the server. +This can be any python file name. This will run a fastapi server. You need to paste a loopback url such as `http://127.0.0.1:8100` inside `Register New Session` in the UI. Below shows, how the basic backend architecture looks like on how Playground communicates with the server. ![](/assets/agent_server.png) diff --git a/premsql/agents/__init__.py b/premsql/agents/__init__.py index 58bd667..07a2fd3 100644 --- a/premsql/agents/__init__.py +++ b/premsql/agents/__init__.py @@ -1,4 +1,11 @@ -from premsql.agents.baseline.main import BaseLineAgent -from premsql.agents.memory import AgentInteractionMemory +from importlib import import_module -__all__ = ["BaseLineAgent", "AgentInteractionMemory"] \ No newline at end of file +__all__ = ["BaseLineAgent", "AgentInteractionMemory"] + + +def __getattr__(name): + if name == "BaseLineAgent": + return import_module("premsql.agents.baseline.main").BaseLineAgent + if name == "AgentInteractionMemory": + return import_module("premsql.agents.memory").AgentInteractionMemory + raise AttributeError(f"module 'premsql.agents' has no attribute {name!r}") diff --git a/premsql/agents/base.py b/premsql/agents/base.py index 693d292..39cb2e6 100644 --- a/premsql/agents/base.py +++ b/premsql/agents/base.py @@ -158,13 +158,14 @@ def __call__( question: str, input_dataframe: Optional[dict] = None, server_mode: Optional[bool] = False, + generate_plot_image: Optional[bool] = True, ) -> Union[ExitWorkerOutput, AgentOutput]: if server_mode: kwargs = self.route_worker_kwargs.get("plot", None) kwargs = ( - {"plot_image": False} + {"plot_image": generate_plot_image} if kwargs is None - else {**kwargs, "plot_image": False} + else {**kwargs, "plot_image": generate_plot_image} ) self.route_worker_kwargs["plot"] = kwargs diff --git a/premsql/agents/baseline/prompts.py b/premsql/agents/baseline/prompts.py index 48a1f9c..36e509f 100644 --- a/premsql/agents/baseline/prompts.py +++ b/premsql/agents/baseline/prompts.py @@ -129,6 +129,7 @@ a single line (string format). - Make sure the column names are correct and exists in the table - For column names which has a space with it, make sure you have put `` in that column name +- Only generate read-only SQL queries. Never generate INSERT, UPDATE, DELETE, DROP, ALTER, or PRAGMA statements. # Database and Table Schema: {schemas} @@ -145,6 +146,7 @@ a single line (string format). - Make sure the column names are correct and exists in the table - For column names which has a space with it, make sure you have put `` in that column name +- Only generate read-only SQL queries. Never generate INSERT, UPDATE, DELETE, DROP, ALTER, or PRAGMA statements. # Database and Table Schema: {schemas} diff --git a/premsql/agents/baseline/workers/followup.py b/premsql/agents/baseline/workers/followup.py index 13f58dc..2fc5507 100644 --- a/premsql/agents/baseline/workers/followup.py +++ b/premsql/agents/baseline/workers/followup.py @@ -7,6 +7,7 @@ from premsql.agents.base import WorkerBase from premsql.agents.baseline.prompts import BASELINE_FOLLOWUP_WORKER_PROMPT from premsql.agents.models import ExitWorkerOutput, FollowupWorkerOutput +from premsql.security import ensure_expected_keys_only, parse_structured_output logger = setup_console_logger("[BASELINE-FOLLOWUP-WORKER]") @@ -78,7 +79,13 @@ def run( max_new_tokens=max_new_tokens, postprocess=False, ) - result = eval(result.replace("null", "None")) + result = ensure_expected_keys_only( + parse_structured_output( + result, + expected_keys={"alternate_decision", "suggestion"}, + ), + expected_keys={"alternate_decision", "suggestion"}, + ) error_from_model = None assert "alternate_decision" in result assert "suggestion" in result diff --git a/premsql/agents/baseline/workers/plotter.py b/premsql/agents/baseline/workers/plotter.py index 1b3a133..3ef4322 100644 --- a/premsql/agents/baseline/workers/plotter.py +++ b/premsql/agents/baseline/workers/plotter.py @@ -8,6 +8,7 @@ from premsql.agents.baseline.prompts import BASELINE_CHART_WORKER_PROMPT_TEMPLATE from premsql.agents.tools.plot.base import BasePlotTool from premsql.agents.utils import convert_df_to_dict +from premsql.security import ensure_expected_keys_only, parse_structured_output logger = setup_console_logger("[PLOT-WORKER]") @@ -39,8 +40,12 @@ def run( max_new_tokens=max_new_tokens, postprocess=False, ) - to_plot = to_plot.replace("null", "None") - plot_config = eval(to_plot) + plot_config = ensure_expected_keys_only( + parse_structured_output( + to_plot, expected_keys={"x", "y", "plot_type"} + ), + expected_keys={"x", "y", "plot_type"}, + ) fig = self.plot_tool.run(data=input_dataframe, plot_config=plot_config) logger.info(f"Plot config: {plot_config}") @@ -48,7 +53,7 @@ def run( output = self.plot_tool.convert_image_to_base64( self.plot_tool.convert_plot_to_image(fig=fig) ) - logger.info("Done base64 conversion") + logger.info("Plot image generated successfully") else: output = None @@ -69,6 +74,7 @@ def run( except Exception as e: error_message = f"Error during plot generation: {str(e)}" + logger.error(error_message) return ChartPlotWorkerOutput( question=question, input_dataframe=convert_df_to_dict(input_dataframe), diff --git a/premsql/agents/baseline/workers/text2sql.py b/premsql/agents/baseline/workers/text2sql.py index d7ddd1d..ca84a9f 100644 --- a/premsql/agents/baseline/workers/text2sql.py +++ b/premsql/agents/baseline/workers/text2sql.py @@ -13,6 +13,11 @@ ) from premsql.agents.models import Text2SQLWorkerOutput from premsql.agents.utils import execute_and_render_result +from premsql.security import ( + SecurityValidationError, + ensure_expected_keys_only, + parse_structured_output, +) logger = setup_console_logger("[BASELINE-TEXT2SQL-WORKER]") @@ -27,6 +32,7 @@ def __init__( include_tables: Optional[list] = None, exclude_tables: Optional[list] = None, auto_filter_tables: Optional[bool] = False, + allow_raw_sql: Optional[bool] = False, ): super().__init__( db_connection_uri=db_connection_uri, @@ -39,6 +45,7 @@ def __init__( self.corrector = helper_model self.table_filer_worker = helper_model self.auto_filter_tables = auto_filter_tables + self.allow_raw_sql = allow_raw_sql @staticmethod def show_dataframe(output: Text2SQLWorkerOutput): @@ -64,7 +71,12 @@ def filer_tables_from_schema( try: to_include = [] output = self.corrector.generate({"prompt": prompt}, postprocess=False) - output = eval(output) + output = ensure_expected_keys_only( + parse_structured_output(output, expected_keys={"include"}), + expected_keys={"include"}, + ) + if not isinstance(output["include"], list): + raise SecurityValidationError("Expected 'include' to be a list") for table in all_tables: if table in output["include"]: to_include.append(table) @@ -154,25 +166,44 @@ def run( **kwargs, ) -> Text2SQLWorkerOutput: if question.startswith("`") and question.endswith("`"): + if not self.allow_raw_sql: + return Text2SQLWorkerOutput( + db_connection_uri=self.db_connection_uri, + sql_string=None, + sql_reasoning=None, + input_dataframe=None, + output_dataframe={"data": {}, "columns": []}, + question=question, + error_from_model=( + "Direct SQL execution is disabled. Use natural language prompts instead." + ), + additional_input={ + "additional_knowledge": additional_knowledge, + "fewshot_dict": fewshot_dict, + "temperature": temperature, + "max_new_tokens": max_new_tokens, + **kwargs, + }, + ) result = execute_and_render_result( - db=self.db, sql=question.replace('`', ''), using=render_results_using - ) + db=self.db, sql=question.replace("`", ""), using=render_results_using + ) return Text2SQLWorkerOutput( - db_connection_uri=self.db_connection_uri, - sql_string=question.startswith("`"), - sql_reasoning=None, - input_dataframe=None, - output_dataframe=result["dataframe"], # Truncating to - question=question, - error_from_model=result["error_from_model"], - additional_input={ - "additional_knowledge": additional_knowledge, - "fewshot_dict": fewshot_dict, - "temperature": temperature, - "max_new_tokens": max_new_tokens, - **kwargs, - }, - ) + db_connection_uri=self.db_connection_uri, + sql_string=question.replace("`", ""), + sql_reasoning=None, + input_dataframe=None, + output_dataframe=result["dataframe"], + question=question, + error_from_model=result["error_from_model"], + additional_input={ + "additional_knowledge": additional_knowledge, + "fewshot_dict": fewshot_dict, + "temperature": temperature, + "max_new_tokens": max_new_tokens, + **kwargs, + }, + ) prompt = self._create_prompt( diff --git a/premsql/agents/memory.py b/premsql/agents/memory.py index 95cb2b3..12c87aa 100644 --- a/premsql/agents/memory.py +++ b/premsql/agents/memory.py @@ -1,12 +1,11 @@ import os -import tempfile import sqlite3 -from platformdirs import user_cache_dir from typing import List, Literal, Optional from premsql.logger import setup_console_logger from premsql.agents.models import ExitWorkerOutput from premsql.agents.utils import convert_exit_output_to_agent_output +from premsql.security import quote_sqlite_identifier, validate_session_name logger = setup_console_logger("[PIPELINE-MEMORY]") @@ -14,7 +13,7 @@ class AgentInteractionMemory: def __init__(self, session_name: str, db_path: Optional[str] = None): - self.session_name = session_name + self.session_name = validate_session_name(session_name) self.db_path = db_path or os.path.join( os.getcwd(), "premsql", "premsql_pipeline_memory.db" ) @@ -24,6 +23,10 @@ def __init__(self, session_name: str, db_path: Optional[str] = None): self.conn = sqlite3.connect(self.db_path) self.create_table_if_not_exists() + @property + def table_name(self) -> str: + return quote_sqlite_identifier(self.session_name) + def list_sessions(self) -> List[str]: cursor = self.conn.cursor() @@ -35,7 +38,7 @@ def create_table_if_not_exists(self): cursor = self.conn.cursor() cursor.execute( f""" - CREATE TABLE IF NOT EXISTS {self.session_name} ( + CREATE TABLE IF NOT EXISTS {self.table_name} ( message_id INTEGER PRIMARY KEY AUTOINCREMENT, question TEXT, db_connection_uri TEXT, @@ -71,7 +74,7 @@ def get( order: Optional[Literal["DESC", "ASC"]] = "DESC", ) -> List[tuple[int, ExitWorkerOutput]]: cursor = self.conn.cursor() - query = f"SELECT * FROM {self.session_name} ORDER BY message_id {order}" + query = f"SELECT * FROM {self.table_name} ORDER BY message_id {order}" if limit is not None: query += " LIMIT ?" cursor.execute(query, (limit,)) @@ -85,7 +88,7 @@ def get( def get_latest_message_id(self) -> Optional[int]: cursor = self.conn.cursor() - query = f"SELECT message_id FROM {self.session_name} ORDER BY message_id DESC LIMIT 1" + query = f"SELECT message_id FROM {self.table_name} ORDER BY message_id DESC LIMIT 1" cursor.execute(query) row = cursor.fetchone() return row[0] if row else None @@ -93,9 +96,10 @@ def get_latest_message_id(self) -> Optional[int]: def generate_messages_from_session( self, session_name: str, limit: int = 100, server_mode: bool=False ): + validated_session_name = quote_sqlite_identifier(validate_session_name(session_name)) cursor = self.conn.cursor() - query = f"SELECT * FROM {session_name} ORDER BY message_id ASC LIMIT {limit}" - cursor.execute(query) + query = f"SELECT * FROM {validated_session_name} ORDER BY message_id ASC LIMIT ?" + cursor.execute(query, (limit,)) rows = cursor.fetchall() for row in rows: yield self._row_to_exit_worker_output(row=row) if server_mode == False else convert_exit_output_to_agent_output( @@ -104,7 +108,7 @@ def generate_messages_from_session( def get_by_message_id(self, message_id: int) -> Optional[dict]: cursor = self.conn.cursor() - query = f"SELECT * FROM {self.session_name} WHERE message_id = ?" + query = f"SELECT * FROM {self.table_name} WHERE message_id = ?" cursor.execute(query, (message_id,)) row = cursor.fetchone() if row is None: @@ -115,7 +119,7 @@ def push(self, output: ExitWorkerOutput): cursor = self.conn.cursor() cursor.execute( f""" - INSERT INTO {self.session_name} ( + INSERT INTO {self.table_name} ( question, db_connection_uri, route_taken, sql_string, sql_reasoning, sql_input_dataframe, sql_output_dataframe, error_from_sql_worker, analysis, analysis_reasoning, analysis_input_dataframe, @@ -133,7 +137,7 @@ def push(self, output: ExitWorkerOutput): def delete_table(self): cursor = self.conn.cursor() try: - cursor.execute(f"DROP TABLE IF EXISTS {self.session_name}") + cursor.execute(f"DROP TABLE IF EXISTS {self.table_name}") self.conn.commit() logger.info(f"Table '{self.session_name}' has been deleted.") except sqlite3.Error as e: @@ -222,7 +226,7 @@ def _serialize_json(self, obj): def clear(self): cursor = self.conn.cursor() - cursor.execute(f"DELETE FROM {self.session_name}") + cursor.execute(f"DELETE FROM {self.table_name}") self.conn.commit() def close(self): @@ -234,7 +238,7 @@ def __del__(self): def delete_table(self): cursor = self.conn.cursor() try: - cursor.execute(f"DROP TABLE IF EXISTS {self.session_name}") + cursor.execute(f"DROP TABLE IF EXISTS {self.table_name}") self.conn.commit() logger.info(f"Table '{self.session_name}' has been deleted.") except sqlite3.Error as e: @@ -249,7 +253,7 @@ def get_latest_dataframe( if not contents: return {} - _, content = contents[0] + content = contents[0]["message"] if decision == "plot" and content.plot_input_dataframe: return content.plot_input_dataframe elif decision == "analyse" and content.analysis_input_dataframe: diff --git a/premsql/agents/models.py b/premsql/agents/models.py index 4f4e470..f24a6fd 100644 --- a/premsql/agents/models.py +++ b/premsql/agents/models.py @@ -133,7 +133,7 @@ class AgentOutput(BaseModel): session_name: str question: str - db_connection_uri: str + db_connection_uri: Optional[str] = None route_taken: Literal["plot", "analyse", "query", "followup"] input_dataframe: Optional[Dict] = None output_dataframe: Optional[Dict] = None diff --git a/premsql/agents/tools/plot/matplotlib_tool.py b/premsql/agents/tools/plot/matplotlib_tool.py index ebc5ff2..a7ab341 100644 --- a/premsql/agents/tools/plot/matplotlib_tool.py +++ b/premsql/agents/tools/plot/matplotlib_tool.py @@ -1,6 +1,8 @@ import io from typing import Callable, Dict +import matplotlib +matplotlib.use('Agg') # Use non-interactive backend for server mode import matplotlib.pyplot as plt import pandas as pd from matplotlib.axes import Axes @@ -54,15 +56,18 @@ def _validate_config(self, df: pd.DataFrame, plot_config: Dict[str, str]) -> Non f"Missing required keys in plot_config: {', '.join(missing_keys)}" ) + if plot_config["plot_type"] not in self.plot_functions: + raise ValueError(f"Unsupported plot type: {plot_config['plot_type']}") + if plot_config["x"] not in df.columns: raise ValueError(f"Column '{plot_config['x']}' not found in DataFrame") + if plot_config["plot_type"] == "histogram": + return + if plot_config["y"] not in df.columns: raise ValueError(f"Column '{plot_config['y']}' not found in DataFrame") - if plot_config["plot_type"] not in self.plot_functions: - raise ValueError(f"Unsupported plot type: {plot_config['plot_type']}") - def _area_plot(self, df: pd.DataFrame, x: str, y: str, ax: Axes) -> None: ax.fill_between(df[x], df[y]) diff --git a/premsql/agents/utils.py b/premsql/agents/utils.py index f799c8b..8d05980 100644 --- a/premsql/agents/utils.py +++ b/premsql/agents/utils.py @@ -5,6 +5,7 @@ from premsql.executors.from_langchain import SQLDatabase from premsql.logger import setup_console_logger from premsql.agents.models import AgentOutput, ExitWorkerOutput +from premsql.security import UnsafeSQLQuery, enforce_read_only_sql logger = setup_console_logger("[PIPELINE-UTILS]") @@ -16,11 +17,16 @@ def convert_df_to_dict(df: pd.DataFrame): def execute_and_render_result( db: SQLDatabase, sql: str, using: Literal["dataframe", "json"] ): - result = db.run_no_throw(command=sql, fetch="cursor") + try: + safe_sql = enforce_read_only_sql(sql) + except UnsafeSQLQuery as exc: + return _render_error(str(exc), sql, using) + + result = db.run_no_throw(command=safe_sql, fetch="cursor") if isinstance(result, str): - return _render_error(result, sql, using) - return _render_data(result, sql, using) + return _render_error(result, safe_sql, using) + return _render_data(result, safe_sql, using) def _render_error(error: str, sql: str, using: str) -> Dict[str, Any]: @@ -80,4 +86,4 @@ def convert_exit_output_to_agent_output(exit_output: ExitWorkerOutput) -> AgentO or exit_output.error_from_plot_worker or exit_output.error_from_followup_worker ), - ) \ No newline at end of file + ) diff --git a/premsql/cli.py b/premsql/cli.py index 36216f8..814c505 100644 --- a/premsql/cli.py +++ b/premsql/cli.py @@ -1,16 +1,100 @@ import os import subprocess import sys +import json +import atexit from pathlib import Path import click +# Load environment variables from .env file +try: + from dotenv import load_dotenv + env_file = Path(__file__).parent.parent / ".env" + if env_file.exists(): + load_dotenv(env_file) +except ImportError: + pass + +# PID file for tracking spawned processes +PID_FILE = Path(user_cache_dir()) / "premsql" / "pids.json" if 'user_cache_dir' in dir() else Path.home() / ".premsql_pids.json" + +def _get_pid_file(): + """Get the PID file path.""" + try: + from platformdirs import user_cache_dir + cache_dir = Path(user_cache_dir()) / "premsql" + except ImportError: + cache_dir = Path.home() / ".premsql_cache" + cache_dir.mkdir(parents=True, exist_ok=True) + return cache_dir / "pids.json" + + +def _save_pid(service_name: str, pid: int): + """Save a PID to the PID file.""" + pid_file = _get_pid_file() + pids = {} + if pid_file.exists(): + try: + pids = json.load(pid_file.open("r")) + except (json.JSONDecodeError, IOError): + pids = {} + pids[service_name] = pid + json.dump(pids, pid_file.open("w"), indent=2) + + +def _load_pids(): + """Load all saved PIDs.""" + pid_file = _get_pid_file() + if not pid_file.exists(): + return {} + try: + return json.load(pid_file.open("r")) + except (json.JSONDecodeError, IOError): + return {} + + +def _clear_pid_file(): + """Clear the PID file.""" + pid_file = _get_pid_file() + if pid_file.exists(): + pid_file.unlink() + + +def _stop_process_by_pid(pid: int, timeout: int = 5) -> bool: + """Stop a process by its PID with timeout.""" + try: + proc = subprocess.Popen(["ps", "-p", str(pid)], stdout=subprocess.PIPE, stderr=subprocess.PIPE) + proc.wait(timeout=1) + if proc.returncode == 0: + # Process exists, send SIGTERM + os.kill(pid, 15) # SIGTERM + # Wait for graceful shutdown + import time + for _ in range(timeout): + try: + os.kill(pid, 0) # Check if process still exists + time.sleep(1) + except OSError: + return True # Process terminated + # Force kill if still running + try: + os.kill(pid, 9) # SIGKILL + except OSError: + pass + return True + except (subprocess.TimeoutExpired, OSError, FileNotFoundError): + pass + return False + + @click.group() @click.version_option() def cli(): """PremSQL CLI to manage API servers and Streamlit app""" pass + @cli.group() def launch(): """Launch PremSQL services""" @@ -23,7 +107,7 @@ def launch_all(): premsql_path = Path(__file__).parent.parent.absolute() env = os.environ.copy() env["PYTHONPATH"] = str(premsql_path) - + # Start API server manage_py_path = premsql_path / "premsql" / "playground" / "backend" / "manage.py" if not manage_py_path.exists(): @@ -40,7 +124,13 @@ def launch_all(): sys.exit(1) click.echo("Starting the PremSQL backend API server...") - subprocess.Popen([sys.executable, str(manage_py_path), "runserver"], env=env) + api_process = subprocess.Popen( + [sys.executable, str(manage_py_path), "runserver"], + env=env, + stdout=subprocess.PIPE, + stderr=subprocess.PIPE, + ) + _save_pid("api_server", api_process.pid) # Launch the streamlit app click.echo("Starting the PremSQL Streamlit app...") @@ -49,13 +139,16 @@ def launch_all(): click.echo(f"Error: main.py not found at {main_py_path}", err=True) sys.exit(1) - cmd = [sys.executable, "-m", "streamlit", "run", str(main_py_path), "--server.maxUploadSize=500"] + cmd = [sys.executable, "-m", "streamlit", "run", str(main_py_path), "--server.maxUploadSize=100"] try: - subprocess.run(cmd, env=env, check=True) + streamlit_process = subprocess.Popen(cmd, env=env) + _save_pid("streamlit", streamlit_process.pid) + streamlit_process.wait() except KeyboardInterrupt: click.echo("Stopping all services...") stop() + @launch.command(name='api') def launch_api(): """Launch only the API server""" @@ -80,30 +173,51 @@ def launch_api(): click.echo("Starting the PremSQL backend API server...") cmd = [sys.executable, str(manage_py_path), "runserver"] try: - subprocess.run(cmd, env=env, check=True) + api_process = subprocess.Popen(cmd, env=env) + _save_pid("api_server", api_process.pid) + api_process.wait() except KeyboardInterrupt: click.echo("API server stopped.") + @cli.command() def stop(): - """Stop all PremSQL services""" + """Stop all PremSQL services using saved PIDs""" click.echo("Stopping all PremSQL services...") - try: - if sys.platform == "win32": - subprocess.run( - ["taskkill", "/F", "/IM", "python.exe", "/FI", "WINDOWTITLE eq premsql*"], - check=True, - ) + pids = _load_pids() + if not pids: + click.echo("No saved PIDs found. Services may not have been started via CLI.") + return + + stopped_count = 0 + for service_name, pid in pids.items(): + click.echo(f"Stopping {service_name} (PID: {pid})...") + if _stop_process_by_pid(pid): + click.echo(f" {service_name} stopped successfully.") + stopped_count += 1 else: - subprocess.run(["pkill", "-f", "manage.py runserver"], check=True) - subprocess.run(["pkill", "-f", "streamlit"], check=True) - click.echo("All services stopped successfully.") - except subprocess.CalledProcessError: - click.echo("No running services found.") - except Exception as e: - click.echo(f"Error stopping services: {e}", err=True) - sys.exit(1) + click.echo(f" {service_name} was not running or could not be stopped.") + + _clear_pid_file() + click.echo(f"Stopped {stopped_count} service(s).") + + +@cli.command() +def status(): + """Show status of PremSQL services""" + pids = _load_pids() + if not pids: + click.echo("No saved PIDs found.") + return + + for service_name, pid in pids.items(): + try: + os.kill(pid, 0) # Check if process exists + click.echo(f"{service_name}: running (PID: {pid})") + except OSError: + click.echo(f"{service_name}: not running (saved PID: {pid})") + if __name__ == "__main__": - cli() \ No newline at end of file + cli() diff --git a/premsql/datasets/base.py b/premsql/datasets/base.py index 087d5da..95caadb 100644 --- a/premsql/datasets/base.py +++ b/premsql/datasets/base.py @@ -59,8 +59,10 @@ def schema_prompt(self, db_path: str) -> str: table_name = table[0] if table_name == "sqlite_sequence": continue + # Use parameterized query instead of f-string for safety cursor.execute( - f"SELECT sql FROM sqlite_master WHERE type='table' AND name='{table_name}';" + "SELECT sql FROM sqlite_master WHERE type='table' AND name=?;", + (table_name,) ) create_table_sql = cursor.fetchone() if create_table_sql: @@ -115,7 +117,18 @@ class SupervisedDatasetForTraining(torch.utils.data.Dataset): @classmethod def load_from_pth(cls, dataset_path: Union[str, Path]): dataset_path = str(dataset_path) - dataset_dict = torch.load(dataset_path) + # Security: Use weights_only=True to prevent malicious pickle deserialization + # This is supported in PyTorch 2.0+ + try: + dataset_dict = torch.load(dataset_path, weights_only=True) + except TypeError: + # PyTorch < 2.0 fallback - warn about security risk + import warnings + warnings.warn( + "Loading dataset with legacy torch.load. " + "Consider upgrading to PyTorch 2.0+ for safer deserialization." + ) + dataset_dict = torch.load(dataset_path) assert "input_ids" in dataset_dict[0], "input_ids is required" assert "labels" in dataset_dict[0], "labels is required" diff --git a/premsql/executors/__init__.py b/premsql/executors/__init__.py index f5bac0e..55d44a0 100644 --- a/premsql/executors/__init__.py +++ b/premsql/executors/__init__.py @@ -1,4 +1,21 @@ -from premsql.executors.from_langchain import ExecutorUsingLangChain -from premsql.executors.from_sqlite import SQLiteExecutor, OptimizedSQLiteExecutor +from importlib import import_module __all__ = ["ExecutorUsingLangChain", "SQLiteExecutor", "OptimizedSQLiteExecutor"] + + +def __getattr__(name): + mapping = { + "ExecutorUsingLangChain": ( + "premsql.executors.from_langchain", + "ExecutorUsingLangChain", + ), + "SQLiteExecutor": ("premsql.executors.from_sqlite", "SQLiteExecutor"), + "OptimizedSQLiteExecutor": ( + "premsql.executors.from_sqlite", + "OptimizedSQLiteExecutor", + ), + } + if name not in mapping: + raise AttributeError(f"module 'premsql.executors' has no attribute {name!r}") + module_name, attr_name = mapping[name] + return getattr(import_module(module_name), attr_name) diff --git a/premsql/executors/from_langchain.py b/premsql/executors/from_langchain.py index b0e5919..637760c 100644 --- a/premsql/executors/from_langchain.py +++ b/premsql/executors/from_langchain.py @@ -4,6 +4,7 @@ from langchain_community.utilities.sql_database import SQLDatabase from premsql.executors.base import BaseExecutor +from premsql.security import enforce_read_only_sql from premsql.utils import convert_sqlite_path_to_dsn @@ -17,8 +18,9 @@ def execute_sql(self, sql: str, dsn_or_db_path: Union[str, SQLDatabase]) -> dict else: db = dsn_or_db_path + safe_sql = enforce_read_only_sql(sql) start_time = time.time() - response = db.run_no_throw(sql) + response = db.run_no_throw(safe_sql) end_time = time.time() error = response if response.startswith("Error") else None diff --git a/premsql/executors/from_sqlite.py b/premsql/executors/from_sqlite.py index c3b2cda..1fda86d 100644 --- a/premsql/executors/from_sqlite.py +++ b/premsql/executors/from_sqlite.py @@ -5,6 +5,7 @@ from typing import Any, Dict, Generator from premsql.executors.base import BaseExecutor from premsql.logger import setup_console_logger +from premsql.security import enforce_read_only_sql class OptimizedSQLiteExecutor(BaseExecutor): @@ -31,15 +32,16 @@ def get_connection(self, db_path: str) -> Generator[sqlite3.Connection, None, No def execute_sql(self, sql: str, dsn_or_db_path: str) -> Dict[str, Any]: start_time = time.time() try: + safe_sql = enforce_read_only_sql(sql) with self.get_connection(dsn_or_db_path) as conn: cursor = conn.cursor() - cursor.execute("EXPLAIN QUERY PLAN " + sql) + cursor.execute("EXPLAIN QUERY PLAN " + safe_sql) query_plan = cursor.fetchall() if any("SCAN TABLE" in str(row) for row in query_plan): self.logger.warn("Warning: Full table scan detected. Consider adding an index.") - cursor.execute(sql) + cursor.execute(safe_sql) result = [dict(row) for row in cursor.fetchall()] error = None except sqlite3.Error as e: @@ -98,7 +100,8 @@ def execute_sql(self, sql: str, dsn_or_db_path: str) -> dict: start_time = time.time() try: - cursor.execute(sql) + safe_sql = enforce_read_only_sql(sql) + cursor.execute(safe_sql) result = cursor.fetchall() error = None except Exception as e: @@ -114,4 +117,4 @@ def execute_sql(self, sql: str, dsn_or_db_path: str) -> dict: "error": error, "execution_time": end_time - start_time, } - return result \ No newline at end of file + return result diff --git a/premsql/generators/__init__.py b/premsql/generators/__init__.py index 485d701..72f3dd4 100644 --- a/premsql/generators/__init__.py +++ b/premsql/generators/__init__.py @@ -1,7 +1,30 @@ -from premsql.generators.huggingface import Text2SQLGeneratorHF -from premsql.generators.openai import Text2SQLGeneratorOpenAI -from premsql.generators.premai import Text2SQLGeneratorPremAI -from premsql.generators.mlx import Text2SQLGeneratorMLX -from premsql.generators.ollama_model import Text2SQLGeneratorOllama +from importlib import import_module -__all__ = ["Text2SQLGeneratorHF", "Text2SQLGeneratorPremAI", "Text2SQLGeneratorOpenAI", "Text2SQLGeneratorMLX", "Text2SQLGeneratorOllama"] +__all__ = [ + "Text2SQLGeneratorHF", + "Text2SQLGeneratorPremAI", + "Text2SQLGeneratorOpenAI", + "Text2SQLGeneratorMLX", + "Text2SQLGeneratorOllama", + "Text2SQLGeneratorVLLM", + "Text2SQLGeneratorOpenAICompatible", +] + + +def __getattr__(name): + mapping = { + "Text2SQLGeneratorHF": ("premsql.generators.huggingface", "Text2SQLGeneratorHF"), + "Text2SQLGeneratorPremAI": ("premsql.generators.premai", "Text2SQLGeneratorPremAI"), + "Text2SQLGeneratorOpenAI": ("premsql.generators.openai", "Text2SQLGeneratorOpenAI"), + "Text2SQLGeneratorMLX": ("premsql.generators.mlx", "Text2SQLGeneratorMLX"), + "Text2SQLGeneratorOllama": ("premsql.generators.ollama_model", "Text2SQLGeneratorOllama"), + "Text2SQLGeneratorVLLM": ("premsql.generators.vllm", "Text2SQLGeneratorVLLM"), + "Text2SQLGeneratorOpenAICompatible": ( + "premsql.generators.openai_compatible", + "Text2SQLGeneratorOpenAICompatible", + ), + } + if name not in mapping: + raise AttributeError(f"module 'premsql.generators' has no attribute {name!r}") + module_name, attr_name = mapping[name] + return getattr(import_module(module_name), attr_name) \ No newline at end of file diff --git a/premsql/generators/base.py b/premsql/generators/base.py index 6f8c3e9..1625c34 100644 --- a/premsql/generators/base.py +++ b/premsql/generators/base.py @@ -4,10 +4,14 @@ from pathlib import Path from typing import Optional -import sqlparse from tqdm.auto import tqdm from platformdirs import user_cache_dir +try: + import sqlparse +except ImportError: # pragma: no cover - fallback exercised in minimal envs + sqlparse = None + from premsql.evaluator.base import BaseExecutor from premsql.logger import setup_console_logger from premsql.prompts import ERROR_HANDLING_PROMPT @@ -98,9 +102,6 @@ def execution_guided_decoding( def postprocess(self, output_string: str): sql_start_keywords = [ r"\bSELECT\b", - r"\bINSERT\b", - r"\bUPDATE\b", - r"\bDELETE\b", r"\bWITH\b", ] @@ -112,7 +113,10 @@ def postprocess(self, output_string: str): else: sql_statement = output_string - return sqlparse.format(sql_statement.split("# SQL:")[-1].strip()) + sql_statement = sql_statement.split("# SQL:")[-1].strip() + if sqlparse is not None: + return sqlparse.format(sql_statement) + return sql_statement def load_results_from_folder(self): item_names = [item.name for item in self.experiment_path.iterdir()] diff --git a/premsql/generators/openai.py b/premsql/generators/openai.py index 2b93e50..a102fba 100644 --- a/premsql/generators/openai.py +++ b/premsql/generators/openai.py @@ -1,3 +1,14 @@ +""" +OpenAI Generator for PremSQL + +This generator uses the official OpenAI API for text-to-SQL generation. +Supports GPT-4, GPT-3.5, and other OpenAI models. + +For OpenAI-compatible services (vLLM, LM Studio, etc.), use: +- Text2SQLGeneratorVLLM for vLLM deployments +- Text2SQLGeneratorOpenAICompatible for any OpenAI-compatible API +""" + import os from typing import Optional @@ -6,10 +17,19 @@ try: from openai import OpenAI except ImportError: - raise ImportError("Module openai is not installed") + raise ImportError("Module openai is not installed. Run: pip install openai") class Text2SQLGeneratorOpenAI(Text2SQLGeneratorBase): + """ + Generator using official OpenAI API. + + Environment Variables: + OPENAI_API_KEY: Your OpenAI API key (required) + OPENAI_BASE_URL: Optional custom base URL (for proxies) + OPENAI_ORG_ID: Optional organization ID + """ + def __init__( self, model_name: str, @@ -17,8 +37,34 @@ def __init__( type: str, experiment_folder: Optional[str] = None, openai_api_key: Optional[str] = None, + base_url: Optional[str] = None, + organization: Optional[str] = None, + **kwargs ): + """ + Initialize OpenAI generator. + + Args: + model_name: OpenAI model name (e.g., "gpt-4o-mini", "gpt-4o", "gpt-3.5-turbo") + experiment_name: Name for this experiment + type: Experiment type (e.g., "test", "train") + experiment_folder: Custom folder for experiment results + openai_api_key: OpenAI API key (or set OPENAI_API_KEY env var) + base_url: Custom base URL (optional, for proxies) + organization: OpenAI organization ID (optional) + **kwargs: Additional arguments + """ self._api_key = openai_api_key or os.environ.get("OPENAI_API_KEY") + self._base_url = base_url or os.environ.get("OPENAI_BASE_URL") + self._organization = organization or os.environ.get("OPENAI_ORG_ID") + self._kwargs = kwargs + + if not self._api_key: + raise ValueError( + "OpenAI API key is required. Set OPENAI_API_KEY environment variable " + "or pass openai_api_key parameter." + ) + self.model_name = model_name super().__init__( experiment_folder=experiment_folder, @@ -28,11 +74,17 @@ def __init__( @property def load_client(self): - client = OpenAI(api_key=self._api_key) - return client + """Load OpenAI client""" + client_kwargs = {"api_key": self._api_key} + if self._base_url: + client_kwargs["base_url"] = self._base_url + if self._organization: + client_kwargs["organization"] = self._organization + return OpenAI(**client_kwargs) @property def load_tokenizer(self): + """Tokenizer not needed for OpenAI API""" pass @property @@ -47,19 +99,26 @@ def generate( postprocess: Optional[bool] = True, **kwargs ) -> str: + """ + Generate SQL using OpenAI API. + + Args: + data_blob: Contains 'prompt' key with the input prompt + temperature: Sampling temperature (0.0 = deterministic) + max_new_tokens: Maximum tokens to generate + postprocess: Whether to postprocess output as SQL + **kwargs: Additional generation parameters + """ prompt = data_blob["prompt"] - max_tokens = max_new_tokens generation_config = { **kwargs, - **{"temperature": temperature, "max_tokens": max_tokens}, + **{"temperature": temperature, "max_tokens": max_new_tokens}, } - completion = ( - self.client.chat.completions.create( - model=self.model_name, - messages=[{"role": "user", "content": prompt}], - **generation_config - ) - .choices[0] - .message.content - ) - return self.postprocess(output_string=completion) if postprocess else completion + + completion = self.client.chat.completions.create( + model=self.model_name, + messages=[{"role": "user", "content": prompt}], + **generation_config + ).choices[0].message.content + + return self.postprocess(output_string=completion) if postprocess else completion \ No newline at end of file diff --git a/premsql/generators/openai_compatible.py b/premsql/generators/openai_compatible.py new file mode 100644 index 0000000..d8140f0 --- /dev/null +++ b/premsql/generators/openai_compatible.py @@ -0,0 +1,202 @@ +""" +OpenAI-Compatible Generator for PremSQL + +This generator supports any service that provides an OpenAI-compatible API, +including: +- vLLM +- LM Studio +- LocalAI +- Oobabooga (with OpenAI extension) +- Text Generation WebUI +- Any custom deployment with OpenAI-compatible endpoints + +Usage: + from premsql.generators import Text2SQLGeneratorOpenAICompatible + + # For vLLM + generator = Text2SQLGeneratorOpenAICompatible( + model_name="/models/qwen", + base_url="http://localhost:8000/v1", + experiment_name="text2sql_custom", + type="test" + ) + + # For LM Studio + generator = Text2SQLGeneratorOpenAICompatible( + model_name="local-model", + base_url="http://localhost:1234/v1", + experiment_name="lm_studio", + type="test" + ) + + # For any custom OpenAI-compatible service + generator = Text2SQLGeneratorOpenAICompatible( + model_name="your-model", + base_url="http://your-server:port/v1", + api_key="your-key-if-needed", + experiment_name="custom", + type="test", + extra_body={"custom_param": "value"} # Service-specific params + ) +""" + +import os +from typing import Optional + +from premsql.generators.base import Text2SQLGeneratorBase +from premsql.logger import setup_console_logger + +logger = setup_console_logger(name="[OPENAI-COMPATIBLE-GENERATOR]") + +try: + from openai import OpenAI +except ImportError: + raise ImportError("Module openai is not installed. Run: pip install openai") + + +class Text2SQLGeneratorOpenAICompatible(Text2SQLGeneratorBase): + """ + Universal generator for any OpenAI-compatible API service. + + This generator provides a flexible interface for services that implement + the OpenAI chat completions API, allowing you to use: + - Self-hosted models (vLLM, LM Studio, LocalAI) + - Custom deployments + - Alternative API providers + + The generator supports: + - Custom base_url for any OpenAI-compatible endpoint + - Optional API key (some services don't require authentication) + - extra_body for service-specific parameters + - extra_headers for custom headers + - Model-specific configurations via extra_params + """ + + def __init__( + self, + model_name: str, + experiment_name: str, + type: str, + experiment_folder: Optional[str] = None, + base_url: Optional[str] = None, + api_key: Optional[str] = None, + extra_body: Optional[dict] = None, + extra_headers: Optional[dict] = None, + default_params: Optional[dict] = None, + **kwargs + ): + """ + Initialize OpenAI-compatible generator. + + Args: + model_name: Model identifier as expected by the service + experiment_name: Name for this experiment + type: Experiment type (e.g., "test", "train") + experiment_folder: Custom folder for experiment results + base_url: API endpoint URL (e.g., http://localhost:8000/v1) + api_key: API key (optional for local deployments) + extra_body: Additional body parameters for API requests + (e.g., {"chat_template_kwargs": {"enable_thinking": False}}) + extra_headers: Additional headers for API requests + default_params: Default generation parameters (temperature, etc.) + **kwargs: Additional arguments passed to base class + """ + # Get configuration from environment if not provided + self._base_url = base_url or os.environ.get("OPENAI_COMPATIBLE_BASE_URL") + self._api_key = api_key or os.environ.get("OPENAI_COMPATIBLE_API_KEY") or "compatible-dummy-key" + self._extra_body = extra_body or {} + self._extra_headers = extra_headers or {} + self._default_params = default_params or {} + self._kwargs = kwargs + + self.model_name = model_name + super().__init__( + experiment_folder=experiment_folder, + experiment_name=experiment_name, + type=type, + ) + + if self._base_url: + logger.info(f"OpenAI-compatible generator initialized: {self._base_url}") + else: + logger.warning( + "No base_url provided. Will use OpenAI's default API. " + "Set base_url or OPENAI_COMPATIBLE_BASE_URL environment variable." + ) + + @property + def load_client(self): + """Load OpenAI client with custom configuration""" + client_kwargs = {"api_key": self._api_key} + if self._base_url: + client_kwargs["base_url"] = self._base_url + if self._extra_headers: + client_kwargs["default_headers"] = self._extra_headers + return OpenAI(**client_kwargs) + + @property + def load_tokenizer(self): + """Tokenizer not needed for API-based generators""" + pass + + @property + def model_name_or_path(self): + return self.model_name + + def generate( + self, + data_blob: dict, + temperature: Optional[float] = 0.0, + max_new_tokens: Optional[int] = 256, + postprocess: Optional[bool] = True, + **kwargs + ) -> str: + """ + Generate SQL using the OpenAI-compatible API. + + Args: + data_blob: Contains 'prompt' key with the input prompt + temperature: Sampling temperature (0.0 = deterministic) + max_new_tokens: Maximum tokens to generate + postprocess: Whether to postprocess output as SQL + **kwargs: Additional generation parameters + """ + prompt = data_blob["prompt"] + + # Merge default params with call-specific params + generation_config = { + "temperature": temperature, + "max_tokens": max_new_tokens, + **self._default_params, + **kwargs, + } + + # Prepare API call options + api_options = { + "model": self.model_name, + "messages": [{"role": "user", "content": prompt}], + **generation_config, + } + + # Add extra_body if provided + if self._extra_body: + api_options["extra_body"] = self._extra_body + + try: + response = self.client.chat.completions.create(**api_options) + output = response.choices[0].message.content + except Exception as e: + logger.error(f"OpenAI-compatible generation error: {e}") + raise + + return self.postprocess(output_string=output) if postprocess else output + + def set_extra_body(self, extra_body: dict): + """Update extra_body parameters dynamically""" + self._extra_body.update(extra_body) + logger.info(f"Updated extra_body: {self._extra_body}") + + def set_default_params(self, params: dict): + """Update default generation parameters""" + self._default_params.update(params) + logger.info(f"Updated default_params: {self._default_params}") \ No newline at end of file diff --git a/premsql/generators/vllm.py b/premsql/generators/vllm.py new file mode 100644 index 0000000..ff92199 --- /dev/null +++ b/premsql/generators/vllm.py @@ -0,0 +1,167 @@ +""" +vLLM Generator for PremSQL + +vLLM provides an OpenAI-compatible API server, making it easy to integrate +with PremSQL. This generator handles vLLM-specific configurations like +disable_thinking for Qwen3 models. + +Usage: + # Start vLLM server first: + vllm serve /models/qwen --port 8000 + + # Then use in PremSQL: + from premsql.generators import Text2SQLGeneratorVLLM + + generator = Text2SQLGeneratorVLLM( + model_name="/models/qwen", + base_url="http://localhost:8000/v1", + experiment_name="text2sql_vllm", + type="test" + ) +""" + +import os +from typing import Optional + +from premsql.generators.base import Text2SQLGeneratorBase +from premsql.logger import setup_console_logger + +logger = setup_console_logger(name="[VLLM-GENERATOR]") + +try: + from openai import OpenAI +except ImportError: + raise ImportError("Module openai is not installed. Run: pip install openai") + + +class Text2SQLGeneratorVLLM(Text2SQLGeneratorBase): + """ + Generator for vLLM deployed models. + + vLLM is a high-performance LLM inference server that provides an + OpenAI-compatible API. This generator supports: + + - Any model deployed via vLLM + - Qwen3 thinking mode control (disable_thinking) + - Custom base_url for remote vLLM servers + + Environment Variables: + VLLM_BASE_URL: Base URL for vLLM server (e.g., http://localhost:8000/v1) + VLLM_MODEL_NAME: Model name/path as configured in vLLM + VLLM_API_KEY: Optional API key (vLLM doesn't require real key by default) + """ + + def __init__( + self, + model_name: str, + experiment_name: str, + type: str, + experiment_folder: Optional[str] = None, + base_url: Optional[str] = None, + api_key: Optional[str] = None, + disable_thinking: Optional[bool] = None, + extra_params: Optional[dict] = None, + **kwargs + ): + """ + Initialize vLLM generator. + + Args: + model_name: Model name or path as configured in vLLM serve command + experiment_name: Name for this experiment + type: Experiment type (e.g., "test", "train") + experiment_folder: Custom folder for experiment results + base_url: vLLM server URL (e.g., http://localhost:8000/v1) + api_key: API key (optional, vLLM doesn't require real key) + disable_thinking: Disable Qwen3 thinking mode (auto-detect for Qwen models) + extra_params: Additional parameters to pass to vLLM API + """ + self._base_url = base_url or os.environ.get("VLLM_BASE_URL") + self._api_key = api_key or os.environ.get("VLLM_API_KEY") or "vllm-dummy-key" + self._extra_params = extra_params or {} + self._kwargs = kwargs + + # Auto-detect if thinking mode should be disabled for Qwen models + if disable_thinking is None: + disable_thinking = self._is_qwen_model(model_name) + self._disable_thinking = disable_thinking + + self.model_name = model_name + super().__init__( + experiment_folder=experiment_folder, + experiment_name=experiment_name, + type=type, + ) + + logger.info(f"vLLM generator initialized with base_url: {self._base_url}") + if self._disable_thinking: + logger.info("Qwen thinking mode disabled") + + def _is_qwen_model(self, model_name: str) -> bool: + """Auto-detect if model is Qwen3 which needs thinking mode disabled""" + qwen_patterns = ["qwen", "Qwen", "QWEN"] + return any(pattern.lower() in model_name.lower() for pattern in qwen_patterns) + + @property + def load_client(self): + """Load OpenAI client configured for vLLM""" + if not self._base_url: + raise ValueError( + "vLLM base_url is required. Set VLLM_BASE_URL environment variable " + "or pass base_url parameter." + ) + return OpenAI(api_key=self._api_key, base_url=self._base_url) + + @property + def load_tokenizer(self): + """vLLM doesn't need tokenizer, handled server-side""" + pass + + @property + def model_name_or_path(self): + return self.model_name + + def generate( + self, + data_blob: dict, + temperature: Optional[float] = 0.0, + max_new_tokens: Optional[int] = 256, + postprocess: Optional[bool] = True, + **kwargs + ) -> str: + """ + Generate SQL using vLLM. + + Args: + data_blob: Contains 'prompt' key with the input prompt + temperature: Sampling temperature + max_new_tokens: Maximum tokens to generate + postprocess: Whether to postprocess output as SQL + **kwargs: Additional generation parameters + """ + prompt = data_blob["prompt"] + generation_config = { + "temperature": temperature, + "max_tokens": max_new_tokens, + **self._extra_params, + **kwargs, + } + + # For Qwen3 models, disable thinking mode via extra_body + extra_body = None + if self._disable_thinking: + extra_body = {"chat_template_kwargs": {"enable_thinking": False}} + + try: + response = self.client.chat.completions.create( + model=self.model_name, + messages=[{"role": "user", "content": prompt}], + extra_body=extra_body, + **generation_config + ) + output = response.choices[0].message.content + except Exception as e: + logger.error(f"vLLM generation error: {e}") + raise + + return self.postprocess(output_string=output) if postprocess else output \ No newline at end of file diff --git a/premsql/logger.py b/premsql/logger.py index 606aaf8..a79704d 100644 --- a/premsql/logger.py +++ b/premsql/logger.py @@ -7,11 +7,12 @@ def setup_console_logger(name, level=logging.INFO): "%(asctime)s - %(name)s - %(levelname)s - %(message)s" ) - console_handler = logging.StreamHandler() - console_handler.setFormatter(formatter) - logger = logging.getLogger(name) logger.setLevel(level) - logger.addHandler(console_handler) + if not logger.handlers: + console_handler = logging.StreamHandler() + console_handler.setFormatter(formatter) + logger.addHandler(console_handler) + logger.propagate = False return logger diff --git a/premsql/playground/__init__.py b/premsql/playground/__init__.py index f887599..4a63c2d 100644 --- a/premsql/playground/__init__.py +++ b/premsql/playground/__init__.py @@ -1,5 +1,21 @@ -from premsql.playground.backend.backend_client import BackendAPIClient -from premsql.playground.inference_server.api_client import InferenceServerAPIClient -from premsql.playground.inference_server.service import AgentServer +from importlib import import_module -__all__ = ["AgentServer", "InferenceServerAPIClient", "BackendAPIClient"] \ No newline at end of file +__all__ = ["AgentServer", "InferenceServerAPIClient", "BackendAPIClient"] + + +def __getattr__(name): + mapping = { + "BackendAPIClient": ( + "premsql.playground.backend.backend_client", + "BackendAPIClient", + ), + "InferenceServerAPIClient": ( + "premsql.playground.inference_server.api_client", + "InferenceServerAPIClient", + ), + "AgentServer": ("premsql.playground.inference_server.service", "AgentServer"), + } + if name not in mapping: + raise AttributeError(f"module 'premsql.playground' has no attribute {name!r}") + module_name, attr_name = mapping[name] + return getattr(import_module(module_name), attr_name) diff --git a/premsql/playground/backend/api/migrations/0002_completions_agent_output.py b/premsql/playground/backend/api/migrations/0002_completions_agent_output.py new file mode 100644 index 0000000..28a426f --- /dev/null +++ b/premsql/playground/backend/api/migrations/0002_completions_agent_output.py @@ -0,0 +1,18 @@ +# Generated by Codex on 2026-04-13 + +from django.db import migrations, models + + +class Migration(migrations.Migration): + + dependencies = [ + ("api", "0001_initial"), + ] + + operations = [ + migrations.AddField( + model_name="completions", + name="agent_output", + field=models.JSONField(blank=True, null=True), + ), + ] diff --git a/premsql/playground/backend/api/migrations/0003_alter_session_fields.py b/premsql/playground/backend/api/migrations/0003_alter_session_fields.py new file mode 100644 index 0000000..4cec096 --- /dev/null +++ b/premsql/playground/backend/api/migrations/0003_alter_session_fields.py @@ -0,0 +1,23 @@ +# Generated by Codex on 2026-04-13 + +from django.db import migrations, models + + +class Migration(migrations.Migration): + + dependencies = [ + ("api", "0002_completions_agent_output"), + ] + + operations = [ + migrations.AlterField( + model_name="session", + name="db_connection_uri", + field=models.CharField(max_length=255), + ), + migrations.AlterField( + model_name="session", + name="session_db_path", + field=models.CharField(blank=True, default="", max_length=255), + ), + ] diff --git a/premsql/playground/backend/api/models.py b/premsql/playground/backend/api/models.py index 8b34c80..e0c9a66 100644 --- a/premsql/playground/backend/api/models.py +++ b/premsql/playground/backend/api/models.py @@ -3,11 +3,11 @@ class Session(models.Model): session_id = models.AutoField(primary_key=True) - db_connection_uri = models.URLField() + db_connection_uri = models.CharField(max_length=255) session_name = models.CharField(max_length=255, unique=True) created_at = models.DateTimeField(auto_now_add=True) base_url = models.URLField() - session_db_path = models.CharField(max_length=255) + session_db_path = models.CharField(max_length=255, blank=True, default="") class Meta: ordering = ["created_at"] @@ -22,6 +22,7 @@ class Completions(models.Model): session_name = models.CharField(max_length=255) created_at = models.DateTimeField() question = models.TextField(blank=True, null=True) + agent_output = models.JSONField(blank=True, null=True) class Meta: ordering = ["-created_at"] diff --git a/premsql/playground/backend/api/pydantic_models.py b/premsql/playground/backend/api/pydantic_models.py index fb14a74..6777ad1 100644 --- a/premsql/playground/backend/api/pydantic_models.py +++ b/premsql/playground/backend/api/pydantic_models.py @@ -1,9 +1,10 @@ from datetime import datetime from typing import List, Literal, Optional -from pydantic import BaseModel, ConfigDict, Field +from pydantic import BaseModel, ConfigDict, Field, field_validator from premsql.agents.models import AgentOutput +from premsql.security import normalize_base_url, validate_session_name # All the Session Models @@ -12,15 +13,19 @@ class SessionCreationRequest(BaseModel): base_url: str = Field(...) model_config = ConfigDict(extra="forbid") + @field_validator("base_url") + @classmethod + def validate_base_url(cls, value: str) -> str: + return normalize_base_url(value) + class SessionCreationResponse(BaseModel): - status_code: Literal[200, 500] = Field(...) + status_code: Literal[200, 401, 500] = Field(...) status: Literal["success", "error"] = Field(...) session_id: Optional[int] = None session_name: Optional[str] = None - db_connection_uri: str = Field(None) - session_db_path: str = Field(None) + # db_connection_uri and session_db_path removed - sensitive, never exposed created_at: Optional[datetime] = None error_message: Optional[str] = None @@ -30,14 +35,13 @@ class SessionSummary(BaseModel): session_name: str created_at: datetime base_url: str - db_connection_uri: str - session_db_path: str + # db_connection_uri and session_db_path removed - sensitive, never exposed model_config = ConfigDict(from_attributes=True) class SessionListResponse(BaseModel): - status_code: Literal[200, 500] + status_code: Literal[200, 401, 404, 500] status: Literal["success", "error"] sessions: Optional[List[SessionSummary]] = None total_count: Optional[int] = None @@ -48,7 +52,7 @@ class SessionListResponse(BaseModel): class SessionDeleteResponse(BaseModel): session_name: str - status_code: Literal[200, 404, 500] + status_code: Literal[200, 401, 404, 500] status: Literal["success", "error"] error_message: Optional[str] = None @@ -57,12 +61,17 @@ class SessionDeleteResponse(BaseModel): class CompletionCreationRequest(BaseModel): - session_name: str - question: str + session_name: str = Field(..., min_length=1, max_length=64) + question: str = Field(..., min_length=1, max_length=4000) + + @field_validator("session_name") + @classmethod + def validate_session_name_value(cls, value: str) -> str: + return validate_session_name(value) class CompletionCreationResponse(BaseModel): - status_code: Literal[200, 500] + status_code: Literal[200, 401, 404, 500] status: Literal["success", "error"] message_id: Optional[int] = None session_name: Optional[str] = None @@ -75,16 +84,19 @@ class CompletionCreationResponse(BaseModel): class CompletionSummary(BaseModel): message_id: int session_name: str - base_url: str + # base_url removed - not needed in chat history created_at: datetime question: Optional[str] = None + message: Optional[AgentOutput] = None model_config = ConfigDict(from_attributes=True) class CompletionListResponse(BaseModel): - status_code: Literal[200, 500] + status_code: Literal[200, 401, 404, 500] status: Literal["success", "error"] completions: Optional[List[CompletionSummary]] = None total_count: Optional[int] = None + page: Optional[int] = None + page_size: Optional[int] = None error_message: Optional[str] = None diff --git a/premsql/playground/backend/api/serializers.py b/premsql/playground/backend/api/serializers.py index 3c524f5..ec71b0c 100644 --- a/premsql/playground/backend/api/serializers.py +++ b/premsql/playground/backend/api/serializers.py @@ -4,7 +4,7 @@ class AgentOutputSerializer(serializers.Serializer): session_name = serializers.CharField() question = serializers.CharField() - db_connection_uri = serializers.CharField() + # db_connection_uri removed - sensitive data, never exposed to clients route_taken = serializers.ChoiceField( choices=["plot", "analyse", "query", "followup"] ) @@ -28,13 +28,12 @@ class SessionCreationRequestSerializer(serializers.Serializer): class SessionCreationResponseSerializer(serializers.Serializer): - status_code = serializers.ChoiceField(choices=[200, 500]) + status_code = serializers.ChoiceField(choices=[200, 401, 500]) status = serializers.ChoiceField(choices=["success", "error"]) session_id = serializers.IntegerField(allow_null=True) session_name = serializers.CharField(allow_null=True) - db_connection_uri = serializers.CharField(allow_null=True) - session_db_path = serializers.CharField(allow_null=True) + # db_connection_uri and session_db_path removed - sensitive, never exposed created_at = serializers.DateTimeField(allow_null=True) error_message = serializers.CharField(allow_null=True) @@ -44,12 +43,11 @@ class SessionSummarySerializer(serializers.Serializer): session_name = serializers.CharField(max_length=255) created_at = serializers.DateTimeField() base_url = serializers.CharField() - db_connection_uri = serializers.CharField() - session_db_path = serializers.CharField() + # db_connection_uri and session_db_path removed - sensitive, never exposed class SessionListResponseSerializer(serializers.Serializer): - status_code = serializers.ChoiceField(choices=[200, 500]) + status_code = serializers.ChoiceField(choices=[200, 401, 404, 500]) status = serializers.ChoiceField(choices=["success", "error"]) sessions = SessionSummarySerializer(many=True, allow_null=True) total_count = serializers.IntegerField(allow_null=True) @@ -60,7 +58,7 @@ class SessionListResponseSerializer(serializers.Serializer): class SessionDeletionResponse(serializers.Serializer): session_name = serializers.CharField(max_length=255) - status_code = serializers.ChoiceField(choices=[200, 404, 500]) + status_code = serializers.ChoiceField(choices=[200, 401, 404, 500]) status = serializers.ChoiceField(choices=["success", "error"]) error_message = serializers.CharField(allow_null=True) @@ -72,11 +70,11 @@ class CompletionCreationRequestSerializer(serializers.Serializer): class CompletionCreationResponseSerializer(serializers.Serializer): - status_code = serializers.ChoiceField(choices=[200, 500]) + status_code = serializers.ChoiceField(choices=[200, 401, 404, 500]) status = serializers.ChoiceField(choices=["success", "error"]) message_id = serializers.IntegerField(allow_null=True) session_name = serializers.CharField(allow_null=True) - message = message = AgentOutputSerializer(allow_null=True) + message = AgentOutputSerializer(allow_null=True) created_at = serializers.DateTimeField(allow_null=True) question = serializers.CharField(allow_null=True) error_message = serializers.CharField(allow_null=True) @@ -85,16 +83,19 @@ class CompletionCreationResponseSerializer(serializers.Serializer): class CompletionSummarySerializer(serializers.Serializer): message_id = serializers.IntegerField() session_name = serializers.CharField() - base_url = serializers.CharField() + # base_url removed - not needed in chat history, already known from session created_at = serializers.DateTimeField() question = serializers.CharField(allow_null=True) + message = AgentOutputSerializer(allow_null=True) class CompletionListResponseSerializer(serializers.Serializer): - status_code = serializers.ChoiceField(choices=[200, 500]) + status_code = serializers.ChoiceField(choices=[200, 401, 404, 500]) status = serializers.ChoiceField(choices=["success", "error"]) completions = CompletionSummarySerializer(many=True, allow_null=True) total_count = serializers.IntegerField(allow_null=True) + page = serializers.IntegerField(allow_null=True) + page_size = serializers.IntegerField(allow_null=True) error_message = serializers.CharField(allow_null=True) diff --git a/premsql/playground/backend/api/services.py b/premsql/playground/backend/api/services.py index 0e696d2..cb044d6 100644 --- a/premsql/playground/backend/api/services.py +++ b/premsql/playground/backend/api/services.py @@ -1,7 +1,5 @@ -import subprocess from typing import Optional -import requests from api.models import Completions, Session from api.pydantic_models import ( CompletionCreationRequest, @@ -20,9 +18,15 @@ from premsql.logger import setup_console_logger from premsql.agents.base import AgentOutput -from premsql.agents.memory import AgentInteractionMemory from premsql.playground import InferenceServerAPIClient -from premsql.playground.backend.api.utils import stop_server_on_port +from premsql.security import ( + SecurityValidationError, + get_api_token, + mask_db_connection_uri, + redact_agent_output_payload, + safe_error_message, + validate_session_name, +) logger = setup_console_logger("[SESSION-MANAGER]") @@ -33,43 +37,65 @@ class SessionManageService: def __init__(self) -> None: - self.client = InferenceServerAPIClient() + self.client = InferenceServerAPIClient(api_token=get_api_token()) def create_session( self, request: SessionCreationRequest ) -> SessionCreationResponse: - response = self.client.get_session_info(base_url=request.base_url) + try: + response = self.client.get_session_info(base_url=request.base_url) + except Exception as exc: + logger.error(safe_error_message(exc, debug_mode=False)) + return SessionCreationResponse( + status_code=500, + status="error", + error_message="Unable to contact the inference server", + ) + if response.get("status") == 500: return SessionCreationResponse( status_code=500, status="error", - error_message="Can not start session, internal server error. Try Again!", + error_message="Unable to start the session", ) try: + session_name = validate_session_name(response["session_name"]) + # If session with same name exists, delete it first + existing_session = Session.objects.filter(session_name=session_name).first() + if existing_session: + Completions.objects.filter(session_name=session_name).delete() + existing_session.delete() + logger.info(f"Deleted existing session: {session_name}") + session = Session.objects.create( - session_name=response["session_name"], - db_connection_uri=response["db_connection_uri"], - created_at=response["created_at"], - base_url=response["base_url"], - session_db_path=response["session_db_path"], + session_name=session_name, + db_connection_uri=mask_db_connection_uri(response.get("db_connection_uri")) + or "***", + base_url=request.base_url, + session_db_path="", ) - logger.info(f"Successfully created session: {response['session_name']}") + logger.info(f"Successfully created session: {session_name}") return SessionCreationResponse( status_code=200, status="success", session_id=session.session_id, session_name=session.session_name, - db_connection_uri=response["db_connection_uri"], - session_db_path=response["session_db_path"], created_at=session.created_at, error_message=None, ) - except Exception as e: + except SecurityValidationError: + return SessionCreationResponse( + status_code=500, + status="error", + error_message="Inference server returned an invalid session identifier", + ) + except Exception as exc: + logger.error(safe_error_message(exc, debug_mode=False)) return SessionCreationResponse( status_code=500, status="error", - error_message=f"Can not start session. {e}", + error_message="Unable to create the session", ) def get_session(self, session_name: str) -> Optional[Session]: @@ -89,8 +115,6 @@ def list_session(self, page: int, page_size: int = 20) -> SessionListResponse: session_name=session.session_name, created_at=session.created_at, base_url=session.base_url, - db_connection_uri=session.db_connection_uri, - session_db_path=session.session_db_path, ) for session in page_obj ] @@ -98,19 +122,20 @@ def list_session(self, page: int, page_size: int = 20) -> SessionListResponse: status="success", status_code=200, sessions=session_summaries, - total_count=len(session_summaries), + total_count=paginator.count, page=page, page_size=page_size, ) - except Exception as e: + except Exception as exc: + logger.error(safe_error_message(exc, debug_mode=False)) return SessionListResponse( status="error", status_code=500, - session_summaries=None, + sessions=None, total_count=0, page=page, page_size=page_size, - error_message=f"Error listing sessions: {e}", + error_message="Unable to list sessions", ) def delete_session(self, session_name: str): @@ -118,23 +143,13 @@ def delete_session(self, session_name: str): with transaction.atomic(): session = Session.objects.get(session_name=session_name) try: - running_port = int(session.base_url.split(":")[1]) - stop_server_on_port(port=running_port) - except Exception as e: - logger.info( - "process killing failed, please shut down inference server manually" - ) - pass + self.client.delete_session(base_url=session.base_url) + except Exception: + logger.warning("Unable to notify inference server during session deletion") - # Proceed with deletion Completions.objects.filter(session_name=session_name).delete() session.delete() logger.info("Deleted all the chats") - agent_memory = AgentInteractionMemory( - session_name=session_name, db_path=session.session_db_path - ) - logger.info("Deleted the session registered inside PremSQL Agent") - agent_memory.delete_table() return SessionDeleteResponse( session_name=session_name, status_code=200, @@ -148,18 +163,19 @@ def delete_session(self, session_name: str): status="error", error_message="Session does not exist", ) - except Exception as e: + except Exception as exc: + logger.error(safe_error_message(exc, debug_mode=False)) return SessionDeleteResponse( session_name=session_name, status_code=500, status="error", - error_message=f"Session does not exist: {e}", + error_message="Unable to delete the session", ) class CompletionService: def __init__(self) -> None: - self.client = InferenceServerAPIClient() + self.client = InferenceServerAPIClient(api_token=get_api_token()) def completion( self, request: CompletionCreationRequest @@ -175,34 +191,38 @@ def completion( ) try: - # Small Hack ;_) - base_url = session.base_url - base_url = f"http://{base_url}" session_inference_response = self.client.post_completion( - base_url=base_url, question=request.question + base_url=session.base_url, question=request.question ) - except Exception as e: - logger.error(f"Unexpected error during completion: {str(e)}") + except Exception as exc: + logger.error(safe_error_message(exc, debug_mode=False)) return CompletionCreationResponse( status_code=500, status="error", session_name=session.session_name, - error_message="An unexpected error occurred", + error_message="Unable to process the completion request", ) try: + message_payload = redact_agent_output_payload( + session_inference_response.get("message") + ) + if message_payload is None: + raise ValueError("Completion response did not include a message payload") + + agent_output = AgentOutput(**message_payload) chat = Completions.objects.create( session=session, session_name=session.session_name, question=request.question, message_id=session_inference_response.get("message_id"), - created_at=session_inference_response.get("message").get("created_at"), + created_at=agent_output.created_at, + agent_output=message_payload, ) logger.info( f"Chat completion created successfully for session: {session.session_name}" ) - agent_output = AgentOutput(**session_inference_response.get("message")) return CompletionCreationResponse( status_code=200, status="success", @@ -213,13 +233,13 @@ def completion( message=agent_output, ) - except Exception as e: - logger.error(f"Error saving completion: {str(e)}") + except Exception as exc: + logger.error(safe_error_message(exc, debug_mode=False)) return CompletionCreationResponse( status_code=500, status="error", session_name=session.session_name, - error_message=f"Completion successful, but failed to save: {e}", + error_message="Completion succeeded but could not be stored", ) def chat_history( @@ -249,9 +269,13 @@ def chat_history( CompletionSummary( message_id=completion.message_id, session_name=completion.session_name, - base_url=completion.session.base_url, created_at=completion.created_at, question=completion.question, + message=( + AgentOutput(**completion.agent_output) + if completion.agent_output is not None + else None + ), ) for completion in page_obj ] @@ -264,7 +288,8 @@ def chat_history( page=page, page_size=page_size, ) - except Exception as e: + except Exception as exc: + logger.error(safe_error_message(exc, debug_mode=False)) return CompletionListResponse( status="error", status_code=500, @@ -272,5 +297,5 @@ def chat_history( total_count=0, page=page, page_size=page_size, - error_message=f"Error fetching chat history: {str(e)}", + error_message="Unable to fetch chat history", ) diff --git a/premsql/playground/backend/api/tests.py b/premsql/playground/backend/api/tests.py index 7ce503c..c5e0bc9 100644 --- a/premsql/playground/backend/api/tests.py +++ b/premsql/playground/backend/api/tests.py @@ -1,3 +1,13 @@ -from django.test import TestCase +from django.test import SimpleTestCase -# Create your tests here. +from api.pydantic_models import SessionCreationRequest +from api.utils import clamp_pagination + + +class SecurityValidationTests(SimpleTestCase): + def test_session_creation_request_rejects_non_loopback_base_url(self): + with self.assertRaises(ValueError): + SessionCreationRequest(base_url="http://example.com:8100") + + def test_clamp_pagination_enforces_bounds(self): + self.assertEqual(clamp_pagination(page=0, page_size=500), (1, 100)) diff --git a/premsql/playground/backend/api/utils.py b/premsql/playground/backend/api/utils.py index 5e09bba..b9a67fe 100644 --- a/premsql/playground/backend/api/utils.py +++ b/premsql/playground/backend/api/utils.py @@ -1,25 +1,31 @@ -import logging import os -import signal -import subprocess +from typing import Optional from premsql.logger import setup_console_logger +from premsql.security import PREMSQL_API_TOKEN_HEADER, get_api_token logger = setup_console_logger("[BACKEND-UTILS]") +# Check if in debug mode +DJANGO_DEBUG = os.environ.get("PREMSQL_DJANGO_DEBUG", "false").lower() == "true" -def stop_server_on_port(port: int): - try: - result = subprocess.run( - ["lsof", "-ti", f":{port}"], capture_output=True, text=True - ) - if result.returncode == 0: - pid = int(result.stdout.strip()) - os.kill(pid, signal.SIGTERM) - logger.info(f"Server running on port {port} (PID {pid}) stopped.") - else: - logger.info(f"No server found running on port {port}") - except subprocess.CalledProcessError: - logger.info(f"No server found running on port {port}") - except ProcessLookupError: - logger.info(f"Process on port {port} no longer exists") + +def is_request_authorized(headers: dict, explicit_token: Optional[str] = None) -> bool: + # In debug mode, skip authentication + if DJANGO_DEBUG: + return True + + required_token = get_api_token(explicit_token) + if not required_token: + return True + return headers.get(PREMSQL_API_TOKEN_HEADER) == required_token + + +def clamp_pagination(page: int, page_size: int, max_page_size: int = 100) -> tuple[int, int]: + if page < 1: + page = 1 + if page_size < 1: + page_size = 1 + if page_size > max_page_size: + page_size = max_page_size + return page, page_size diff --git a/premsql/playground/backend/api/views.py b/premsql/playground/backend/api/views.py index 6b62b22..c2add90 100644 --- a/premsql/playground/backend/api/views.py +++ b/premsql/playground/backend/api/views.py @@ -2,6 +2,7 @@ from drf_yasg import openapi from drf_yasg.utils import swagger_auto_schema +from pydantic import ValidationError as PydanticValidationError from rest_framework import status from rest_framework.decorators import api_view from rest_framework.exceptions import ValidationError @@ -25,10 +26,20 @@ ) from .services import CompletionService, SessionManageService +from .utils import clamp_pagination, is_request_authorized logger = setup_console_logger("[VIEWS]") +def _authorize_or_401(request): + if is_request_authorized(request.headers): + return None + return Response( + {"status": "error", "error_message": "Unauthorized"}, + status=status.HTTP_401_UNAUTHORIZED, + ) + + @swagger_auto_schema( method="post", request_body=SessionCreationRequestSerializer, @@ -40,18 +51,22 @@ ) @api_view(["POST"]) def create_session(request): + unauthorized = _authorize_or_401(request) + if unauthorized is not None: + return unauthorized try: session_request = SessionCreationRequest(**request.data) response = SessionManageService().create_session(request=session_request) - return Response(response.model_dump()) - except json.JSONDecodeError: + return Response(response.model_dump(), status=response.status_code) + except (json.JSONDecodeError, ValueError, PydanticValidationError): return Response( - {"status": "error", "error_message": "Invalid JSON"}, + {"status": "error", "error_message": "Invalid request payload"}, status=status.HTTP_400_BAD_REQUEST, ) - except Exception as e: + except Exception: + logger.exception("Unexpected error while creating a session") return Response( - {"status": "error", "error_message": str(e)}, + {"status": "error", "error_message": "Unable to create session"}, status=status.HTTP_500_INTERNAL_SERVER_ERROR, ) @@ -74,9 +89,17 @@ def create_session(request): ) @api_view(["GET"]) def get_session(request, session_name): + unauthorized = _authorize_or_401(request) + if unauthorized is not None: + return unauthorized session = SessionManageService().get_session(session_name=session_name) if session: - session_summary = SessionSummary.model_validate(session) + session_summary = SessionSummary( + session_id=session.session_id, + session_name=session.session_name, + created_at=session.created_at, + base_url=session.base_url, + ) response = SessionListResponse( status="success", status_code=200, @@ -88,7 +111,7 @@ def get_session(request, session_name): else: response = SessionListResponse( status="error", - status_code=500, + status_code=404, error_message="The requested session does not exist.", ) return Response( @@ -123,10 +146,20 @@ def get_session(request, session_name): ) @api_view(["GET"]) def list_sessions(request): - page = int(request.query_params.get("page", 1)) - page_size = int(request.query_params.get("page_size", 20)) + unauthorized = _authorize_or_401(request) + if unauthorized is not None: + return unauthorized + try: + page = int(request.query_params.get("page", 1)) + page_size = int(request.query_params.get("page_size", 20)) + page, page_size = clamp_pagination(page=page, page_size=page_size) + except ValueError: + return Response( + {"status": "error", "error_message": "Invalid page or page_size parameter"}, + status=status.HTTP_400_BAD_REQUEST, + ) response = SessionManageService().list_session(page=page, page_size=page_size) - return Response(response.model_dump()) + return Response(response.model_dump(), status=response.status_code) @swagger_auto_schema( @@ -159,12 +192,16 @@ def list_sessions(request): ) @api_view(["DELETE"]) def delete_session(request, session_name): + unauthorized = _authorize_or_401(request) + if unauthorized is not None: + return unauthorized try: result = SessionManageService().delete_session(session_name=session_name) return Response(result.model_dump(), status=result.status_code) - except Exception as e: + except Exception: + logger.exception("Unexpected error while deleting session") return Response( - {"status": "error", "error_message": str(e)}, + {"status": "error", "error_message": "Unable to delete session"}, status=status.HTTP_500_INTERNAL_SERVER_ERROR, ) @@ -184,25 +221,25 @@ def delete_session(request, session_name): ) @api_view(["POST"]) def create_completion(request): + unauthorized = _authorize_or_401(request) + if unauthorized is not None: + return unauthorized try: completion_request = CompletionCreationRequest(**request.data) response = CompletionService().completion(request=completion_request) return Response( response.model_dump(), - status=( - status.HTTP_200_OK - if response.status == "success" - else status.HTTP_500_INTERNAL_SERVER_ERROR - ), + status=response.status_code, ) - except ValidationError as e: + except (ValidationError, PydanticValidationError, ValueError) as e: return Response( {"status": "error", "error_message": str(e)}, status=status.HTTP_400_BAD_REQUEST, ) - except Exception as e: + except Exception: + logger.exception("Unexpected error while creating completion") return Response( - {"status": "error", "error_message": str(e)}, + {"status": "error", "error_message": "Unable to create completion"}, status=status.HTTP_500_INTERNAL_SERVER_ERROR, ) @@ -241,9 +278,13 @@ def create_completion(request): ) @api_view(["GET"]) def get_chat_history(request, session_name): + unauthorized = _authorize_or_401(request) + if unauthorized is not None: + return unauthorized try: page = int(request.query_params.get("page", 1)) page_size = int(request.query_params.get("page_size", 20)) + page, page_size = clamp_pagination(page=page, page_size=page_size) response = CompletionService().chat_history( session_name=session_name, page=page, page_size=page_size @@ -251,19 +292,16 @@ def get_chat_history(request, session_name): return Response( response.model_dump(), - status=( - status.HTTP_200_OK - if response.status == "success" - else status.HTTP_404_NOT_FOUND - ), + status=response.status_code, ) except ValueError: return Response( {"status": "error", "error_message": "Invalid page or page_size parameter"}, status=status.HTTP_400_BAD_REQUEST, ) - except Exception as e: + except Exception: + logger.exception("Unexpected error while fetching chat history") return Response( - {"status": "error", "error_message": str(e)}, + {"status": "error", "error_message": "Unable to fetch chat history"}, status=status.HTTP_500_INTERNAL_SERVER_ERROR, ) diff --git a/premsql/playground/backend/backend/settings.py b/premsql/playground/backend/backend/settings.py index c132fb2..52344b9 100644 --- a/premsql/playground/backend/backend/settings.py +++ b/premsql/playground/backend/backend/settings.py @@ -10,6 +10,7 @@ https://docs.djangoproject.com/en/5.1/ref/settings/ """ +import os from pathlib import Path # Build paths inside the project like this: BASE_DIR / 'subdir'. @@ -20,12 +21,19 @@ # See https://docs.djangoproject.com/en/5.1/howto/deployment/checklist/ # SECURITY WARNING: keep the secret key used in production secret! -SECRET_KEY = "django-insecure-v3#txach78pic91j!s=ia3w+h@58niky5ozim)j0+6r56m$pmj" +SECRET_KEY = os.environ.get( + "PREMSQL_DJANGO_SECRET_KEY", + "dev-only-premsql-secret-key-change-me", +) # SECURITY WARNING: don't run with debug turned on in production! -DEBUG = True +DEBUG = os.environ.get("PREMSQL_DJANGO_DEBUG", "false").lower() == "true" -ALLOWED_HOSTS = [] +ALLOWED_HOSTS = [ + host.strip() + for host in os.environ.get("PREMSQL_ALLOWED_HOSTS", "127.0.0.1,localhost").split(",") + if host.strip() +] # Application definition diff --git a/premsql/playground/backend/backend/urls.py b/premsql/playground/backend/backend/urls.py index 928f1cf..42aac7f 100644 --- a/premsql/playground/backend/backend/urls.py +++ b/premsql/playground/backend/backend/urls.py @@ -3,6 +3,8 @@ from drf_yasg import openapi from drf_yasg.views import get_schema_view from rest_framework import permissions +from django.conf import settings +from rest_framework.permissions import IsAuthenticated schema_view = get_schema_view( openapi.Info( @@ -12,20 +14,28 @@ contact=openapi.Contact(email="anindyadeep@premai.io"), license=openapi.License(name="MIT"), ), - public=True, - permission_classes=(permissions.AllowAny,), + public=settings.DEBUG, # Only expose schema in DEBUG mode + permission_classes=( + (permissions.AllowAny,) if settings.DEBUG else (IsAuthenticated,) + ), ) urlpatterns = [ path("admin/", admin.site.urls), path("api/", include("api.urls")), - path( - "swagger/", schema_view.without_ui(cache_timeout=0), name="schema-json" - ), - path( - "swagger/", - schema_view.with_ui("swagger", cache_timeout=0), - name="schema-swagger-ui", - ), - path("redoc/", schema_view.with_ui("redoc", cache_timeout=0), name="schema-redoc"), ] + +if settings.DEBUG: + urlpatterns += [ + path( + "swagger/", + schema_view.without_ui(cache_timeout=0), + name="schema-json", + ), + path( + "swagger/", + schema_view.with_ui("swagger", cache_timeout=0), + name="schema-swagger-ui", + ), + path("redoc/", schema_view.with_ui("redoc", cache_timeout=0), name="schema-redoc"), + ] diff --git a/premsql/playground/backend/backend_client.py b/premsql/playground/backend/backend_client.py index 9376397..8ff7c3f 100644 --- a/premsql/playground/backend/backend_client.py +++ b/premsql/playground/backend/backend_client.py @@ -1,14 +1,16 @@ import requests + from premsql.logger import setup_console_logger from premsql.playground.backend.api.pydantic_models import ( - SessionCreationResponse, - SessionDeleteResponse, - SessionListResponse, - SessionCreationRequest, CompletionCreationRequest, CompletionCreationResponse, CompletionListResponse, + SessionCreationRequest, + SessionCreationResponse, + SessionDeleteResponse, + SessionListResponse, ) +from premsql.security import build_auth_headers BASE_URL = "http://127.0.0.1:8000/api" @@ -16,55 +18,44 @@ class BackendAPIClient: - def __init__(self): + def __init__(self, api_token: str | None = None, timeout: int = 180): self.base_url = BASE_URL - self.headers = { - 'accept': 'application/json', - 'Content-Type': 'application/json', - } + self.timeout = timeout + self.headers = build_auth_headers( + { + "accept": "application/json", + "Content-Type": "application/json", + }, + token=api_token, + ) def create_session(self, request: SessionCreationRequest) -> SessionCreationResponse: try: response = requests.post( f"{self.base_url}/session/create", json=request.model_dump(), - headers=self.headers + headers=self.headers, + timeout=self.timeout, ) - response.raise_for_status() # Raises an HTTPError for bad responses - + response.raise_for_status() return SessionCreationResponse(**response.json()) - except requests.RequestException as e: - logger.error(f"Error creating session: {str(e)}") - logger.error(f"Response content: {response.text if 'response' in locals() else 'No response'}") + except requests.RequestException as exc: + logger.error(f"Error creating session: {exc}") return SessionCreationResponse( status="error", - status_code=response.status_code if 'response' in locals() and hasattr(response, 'status_code') else 500, - error_message=f"Failed to create session: {str(e)}" - ) - except ValueError as e: - logger.error(f"Error parsing response: {str(e)}") - logger.error(f"Response content: {response.text if 'response' in locals() else 'No response'}") + status_code=( + response.status_code + if "response" in locals() and hasattr(response, "status_code") + else 500 + ), + error_message="Failed to create session", + ) + except ValueError as exc: + logger.error(f"Error parsing session creation response: {exc}") return SessionCreationResponse( status="error", status_code=500, - error_message=f"Failed to parse server response: {str(e)}" - ) - - except requests.RequestException as e: - logger.error(f"Error creating session: {str(e)}") - logger.error(f"Response content: {response.text}") - return SessionCreationResponse( - status="error", - status_code=response.status_code if hasattr(response, 'status_code') else 500, - error_message=f"Failed to create session: {str(e)}" - ) - except ValueError as e: - logger.error(f"Error parsing response: {str(e)}") - logger.error(f"Response content: {response.text}") - return SessionCreationResponse( - status="error", - status_code=500, - error_message=f"Failed to parse server response: {str(e)}" + error_message="Failed to parse server response", ) def list_sessions(self, page: int = 1, page_size: int = 20) -> SessionListResponse: @@ -72,142 +63,171 @@ def list_sessions(self, page: int = 1, page_size: int = 20) -> SessionListRespon response = requests.get( f"{self.base_url}/session/list/", params={"page": page, "page_size": page_size}, - headers=self.headers + headers=self.headers, + timeout=self.timeout, ) response.raise_for_status() return SessionListResponse(**response.json()) - except requests.RequestException as e: - logger.error(f"Error listing sessions: {str(e)}") - logger.error(f"Response content: {response.text if 'response' in locals() else 'No response'}") + except requests.RequestException as exc: + logger.error(f"Error listing sessions: {exc}") return SessionListResponse( status="error", - status_code=response.status_code if 'response' in locals() and hasattr(response, 'status_code') else 500, - error_message=f"Failed to list sessions: {str(e)}", + status_code=( + response.status_code + if "response" in locals() and hasattr(response, "status_code") + else 500 + ), + error_message="Failed to list sessions", sessions=[], - total_count=0 + total_count=0, + page=page, + page_size=page_size, ) - except ValueError as e: - logger.error(f"Error parsing response: {str(e)}") - logger.error(f"Response content: {response.text if 'response' in locals() else 'No response'}") + except ValueError as exc: + logger.error(f"Error parsing session list response: {exc}") return SessionListResponse( status="error", status_code=500, - error_message=f"Failed to parse server response: {str(e)}", + error_message="Failed to parse server response", sessions=[], - total_count=0 + total_count=0, + page=page, + page_size=page_size, ) def get_session(self, session_name: str) -> SessionListResponse: try: response = requests.get( f"{self.base_url}/session/{session_name}/", - headers=self.headers + headers=self.headers, + timeout=self.timeout, ) response.raise_for_status() return SessionListResponse(**response.json()) - except requests.RequestException as e: - logger.error(f"Error getting session: {str(e)}") - logger.error(f"Response content: {response.text if 'response' in locals() else 'No response'}") + except requests.RequestException as exc: + logger.error(f"Error getting session: {exc}") + status_code = ( + response.status_code + if "response" in locals() and hasattr(response, "status_code") + else 500 + ) return SessionListResponse( status="error", - status_code=response.status_code if 'response' in locals() and hasattr(response, 'status_code') else 500, - error_message=f"Failed to get session: {str(e)}", - name="", - created_at="", - sessions=[] - ) - except (ValueError, KeyError, IndexError) as e: - logger.error(f"Error parsing response: {str(e)}") - logger.error(f"Response content: {response.text if 'response' in locals() else 'No response'}") + status_code=status_code, + error_message="Failed to get session", + sessions=[], + total_count=0, + page=1, + page_size=1, + ) + except ValueError as exc: + logger.error(f"Error parsing session response: {exc}") return SessionListResponse( status="error", status_code=500, - error_message=f"Failed to parse server response: {str(e)}", - name="", - created_at="", - sessions=[] + error_message="Failed to parse server response", + sessions=[], + total_count=0, + page=1, + page_size=1, ) def delete_session(self, session_name: str) -> SessionDeleteResponse: try: response = requests.delete( f"{self.base_url}/session/{session_name}", - headers=self.headers + headers=self.headers, + timeout=self.timeout, ) response.raise_for_status() return SessionDeleteResponse(**response.json()) - except requests.RequestException as e: - logger.error(f"Error deleting session: {str(e)}") - logger.error(f"Response content: {response.text if 'response' in locals() else 'No response'}") + except requests.RequestException as exc: + logger.error(f"Error deleting session: {exc}") return SessionDeleteResponse( + session_name=session_name, status="error", - status_code=response.status_code if 'response' in locals() and hasattr(response, 'status_code') else 500, - error_message=f"Failed to delete session: {str(e)}" - ) - except ValueError as e: - logger.error(f"Error parsing response: {str(e)}") - logger.error(f"Response content: {response.text if 'response' in locals() else 'No response'}") + status_code=( + response.status_code + if "response" in locals() and hasattr(response, "status_code") + else 500 + ), + error_message="Failed to delete session", + ) + except ValueError as exc: + logger.error(f"Error parsing session deletion response: {exc}") return SessionDeleteResponse( + session_name=session_name, status="error", status_code=500, - error_message=f"Failed to parse server response: {str(e)}" + error_message="Failed to parse server response", ) - # Chats - def create_completion(self, request: CompletionCreationRequest) -> CompletionCreationResponse: + def create_completion( + self, request: CompletionCreationRequest + ) -> CompletionCreationResponse: try: response = requests.post( f"{self.base_url}/chat/completion", json=request.model_dump(), - headers=self.headers + headers=self.headers, + timeout=self.timeout, ) response.raise_for_status() return CompletionCreationResponse(**response.json()) - except requests.RequestException as e: - logger.error(f"Error creating completion: {str(e)}") - logger.error(f"Response content: {response.text if 'response' in locals() else 'No response'}") + except requests.RequestException as exc: + logger.error(f"Error creating completion: {exc}") return CompletionCreationResponse( status="error", - status_code=response.status_code if 'response' in locals() and hasattr(response, 'status_code') else 500, - error_message=f"Failed to create completion: {str(e)}", - completion="" - ) - except ValueError as e: - logger.error(f"Error parsing response: {str(e)}") - logger.error(f"Response content: {response.text if 'response' in locals() else 'No response'}") + status_code=( + response.status_code + if "response" in locals() and hasattr(response, "status_code") + else 500 + ), + error_message="Failed to create completion", + ) + except ValueError as exc: + logger.error(f"Error parsing completion response: {exc}") return CompletionCreationResponse( status="error", status_code=500, - error_message=f"Failed to parse server response: {str(e)}", - completion="" + error_message="Failed to parse server response", ) - - def get_chat_history(self, session_name: str, page: int = 1, page_size: int = 20) -> CompletionListResponse: + + def get_chat_history( + self, session_name: str, page: int = 1, page_size: int = 20 + ) -> CompletionListResponse: try: response = requests.get( f"{self.base_url}/chat/history/{session_name}/", params={"page": page, "page_size": page_size}, - headers=self.headers + headers=self.headers, + timeout=self.timeout, ) response.raise_for_status() return CompletionListResponse(**response.json()) - except requests.RequestException as e: - logger.error(f"Error getting chat history: {str(e)}") - logger.error(f"Response content: {response.text if 'response' in locals() else 'No response'}") + except requests.RequestException as exc: + logger.error(f"Error getting chat history: {exc}") return CompletionListResponse( status="error", - status_code=500, - error_message=f"Failed to get chat history: {str(e)}", + status_code=( + response.status_code + if "response" in locals() and hasattr(response, "status_code") + else 500 + ), + error_message="Failed to get chat history", completions=[], - total_count=0 + total_count=0, + page=page, + page_size=page_size, ) - except ValueError as e: - logger.error(f"Error parsing response: {str(e)}") - logger.error(f"Response content: {response.text if 'response' in locals() else 'No response'}") + except ValueError as exc: + logger.error(f"Error parsing chat history response: {exc}") return CompletionListResponse( status="error", status_code=500, - error_message=f"Failed to parse server response: {str(e)}", + error_message="Failed to parse server response", completions=[], - total_count=0 + total_count=0, + page=page, + page_size=page_size, ) diff --git a/premsql/playground/backend/manage.py b/premsql/playground/backend/manage.py index 2606c26..4bb5e09 100644 --- a/premsql/playground/backend/manage.py +++ b/premsql/playground/backend/manage.py @@ -2,6 +2,16 @@ """Django's command-line utility for administrative tasks.""" import os import sys +from pathlib import Path + +# Load environment variables from .env file +try: + from dotenv import load_dotenv + env_file = Path(__file__).parent.parent.parent.parent / ".env" + if env_file.exists(): + load_dotenv(env_file) +except ImportError: + pass def main(): diff --git a/premsql/playground/frontend/components/chat.py b/premsql/playground/frontend/components/chat.py index 84e219e..be65ae6 100644 --- a/premsql/playground/frontend/components/chat.py +++ b/premsql/playground/frontend/components/chat.py @@ -4,7 +4,6 @@ from premsql.playground.inference_server.api_client import InferenceServerAPIClient from premsql.playground.backend.api.pydantic_models import CompletionCreationRequest from premsql.playground.frontend.components.streamlit_plot import StreamlitPlotTool -from premsql.agents.memory import AgentInteractionMemory from premsql.agents.utils import convert_exit_output_to_agent_output from premsql.agents.models import ExitWorkerOutput, AgentOutput from premsql.logger import setup_console_logger @@ -58,27 +57,35 @@ def render_chat_env(self, session_name: str) -> None: ) if session_info.status_code == 500: st.error(f"Failed to render chat History for session: {session_name}") + return + if not session_info.sessions: + st.error(f"Session not found: {session_name}") + return session = session_info.sessions[0] - session_db_path = session.session_db_path base_url = session.base_url - # TODO: Need to understand how can I start the session - - history = AgentInteractionMemory( - session_name=session_name, db_path=session_db_path + history_response = self.backend_client.get_chat_history( + session_name=session_name, + page=1, + page_size=100, ) - - messages = history.generate_messages_from_session(session_name=session_name, server_mode=True) + if history_response.status_code not in (200, 404): + st.error("Failed to load chat history for this session.") + return + messages = history_response.completions or [] if not messages: st.warning("No chat history available for this session.") else: - for message in messages: - with st.chat_message("user"): st.markdown(message.question) + for completion in messages: + with st.chat_message("user"): + st.markdown(completion.question or "") with st.chat_message("assistant"): - self._streamlit_chat_output(message=message) - + if completion.message: + self._streamlit_chat_output(message=completion.message) + else: + st.warning("Message content is unavailable for this chat entry.") + - base_url = f"http://{base_url}" if not base_url.startswith("http://") else base_url is_session_online_status = self.inference_client.is_online(base_url=base_url) if is_session_online_status != 200: st.divider() @@ -98,9 +105,7 @@ def render_chat_env(self, session_name: str) -> None: ) ) if response.status_code == 200: - self._streamlit_chat_output( - message=history.get_by_message_id(message_id=response.message_id) - ) + self._streamlit_chat_output(message=response.message) else: st.error("Something went wrong. Try again") diff --git a/premsql/playground/frontend/components/session.py b/premsql/playground/frontend/components/session.py index 13be3b9..b7d08a0 100644 --- a/premsql/playground/frontend/components/session.py +++ b/premsql/playground/frontend/components/session.py @@ -1,6 +1,8 @@ import streamlit as st -from premsql.playground.backend.backend_client import BackendAPIClient +from pydantic import ValidationError + from premsql.playground.backend.api.pydantic_models import SessionCreationRequest +from premsql.playground.backend.backend_client import BackendAPIClient additional_link_markdown = """ Here are some quick links to get you started with Prem: @@ -10,26 +12,59 @@ - [PremSQL documentation](https://docs.premai.io/premsql/introduction) """ + class SessionComponent: def __init__(self) -> None: self.backend_client = BackendAPIClient() - - def render_list_sessions(self): + + def _refresh_sessions(self): + """Force refresh session list by clearing cache""" + st.rerun() + + def render_session_list_with_actions(self): + """Render session list with delete buttons for each session""" with st.sidebar: - st.sidebar.title("Your Past Sessions") - all_sessions = self.backend_client.list_sessions(page_size=100).sessions - if all_sessions: - all_sessions = [session.session_name for session in all_sessions] + st.title("Sessions") + all_sessions = self.backend_client.list_sessions(page_size=100).sessions or [] + if all_sessions: + # Session selector for chat selected_session = st.selectbox( - label="Your Sessions (refresh if you have created a new one)", - options=all_sessions, + label="Select a session to chat", + options=[session.session_name for session in all_sessions], + key="session_selector", ) + + st.divider() + + # Session management section + st.subheader("Manage Sessions") + + # Create a container for session list with delete buttons + for session in all_sessions: + col1, col2 = st.columns([3, 1]) + with col1: + st.text(session.session_name) + with col2: + if st.button("Delete", key=f"del_{session.session_name}", type="secondary"): + response = self.backend_client.delete_session( + session_name=session.session_name + ) + if response.status_code == 200: + st.success(f"Deleted: {session.session_name}") + st.rerun() + else: + st.error(response.error_message or "Failed to delete") + return selected_session - + else: + st.info("No sessions found. Create a new session below.") + return None + def render_register_session(self): with st.sidebar: - st.sidebar.title("Register new Session") + st.divider() + st.subheader("Create New Session") with st.form( key="session_creation", clear_on_submit=True, @@ -37,37 +72,41 @@ def render_register_session(self): enter_to_submit=False, ): base_url = st.text_input( - label="base_url", - placeholder="the base url in which AgentServer is running" + label="AgentServer URL", + placeholder="http://127.0.0.1:8100", + help="Enter the base URL where AgentServer is running", ) - button = st.form_submit_button(label="Submit") + button = st.form_submit_button(label="Create Session") + if button: - response = self.backend_client.create_session( - request=SessionCreationRequest(base_url=base_url) - ) - if response.status_code == 500: - st.toast(body=st.markdown(response.error_message), icon="❌") + if not base_url: + st.error("Please enter a base URL") + return None + + try: + request = SessionCreationRequest(base_url=base_url) + except ValidationError as exc: + st.error(str(exc)) + return None + + response = self.backend_client.create_session(request=request) + if response.status_code == 200: + st.success(f"Session '{response.session_name}' created successfully") + st.rerun() else: - st.toast(body=f"Session: {response.session_name} created successfully", icon="🥳") + st.error(response.error_message or "Unable to create session") return response - + def render_additional_links(self): with st.sidebar: - with st.container(height=200): + st.divider() + with st.expander("Help & Resources", expanded=False): st.markdown(additional_link_markdown) - + def render_list_sessions(self): + """Legacy method - now combined with actions""" + return self.render_session_list_with_actions() + def render_delete_session_view(self): - with st.sidebar: - with st.expander(label="Delete a session"): - with st.form(key="delete_session", clear_on_submit=True, enter_to_submit=False): - session_name = st.text_input(label="Enter session name") - button = st.form_submit_button(label="Submit") - if button: - all_sessions = self.backend_client.list_sessions(page_size=100).sessions - all_sessions = [session.session_name for session in all_sessions] - if session_name not in all_sessions: - st.error("Session does not exist") - else: - self.backend_client.delete_session(session_name=session_name) - st.success(f"Deleted session: {session_name}. Please refresh") \ No newline at end of file + """Legacy method - delete buttons now integrated in session list""" + pass \ No newline at end of file diff --git a/premsql/playground/frontend/components/streamlit_plot.py b/premsql/playground/frontend/components/streamlit_plot.py index 37279cf..1d4450b 100644 --- a/premsql/playground/frontend/components/streamlit_plot.py +++ b/premsql/playground/frontend/components/streamlit_plot.py @@ -1,9 +1,9 @@ -import traceback from typing import Dict, Any import pandas as pd -import streamlit as st +import streamlit as st from premsql.logger import setup_console_logger from premsql.agents.tools.plot.base import BasePlotTool +from premsql.security import safe_error_message logger = setup_console_logger("[STREAMLIT-TOOL]") @@ -27,12 +27,11 @@ def run(self, data: pd.DataFrame, plot_config: Dict[str, str]) -> Any: st.markdown(f"**{plot_type.capitalize()} Plot: {x} vs {y}**") return self.plot_functions[plot_type](data, x, y) - except Exception as e: - error_msg = f"Error creating plot: {str(e)}" - stack_trace = traceback.format_exc() - logger.error(f"{error_msg}\n{stack_trace}") - logger.error(f"Error creating plot: {str(e)}") - st.error(f"Error creating plot: {str(e)}") + except Exception as exc: + # Use safe error message (whitelist approach) + safe_msg = safe_error_message(exc, debug_mode=False) + logger.error(f"Plot error: {safe_msg}") + st.error(safe_msg) return None def _validate_config(self, df: pd.DataFrame, plot_config: Dict[str, str]) -> None: diff --git a/premsql/playground/frontend/components/uploader.py b/premsql/playground/frontend/components/uploader.py index a6fcddb..9ca1ec2 100644 --- a/premsql/playground/frontend/components/uploader.py +++ b/premsql/playground/frontend/components/uploader.py @@ -21,11 +21,16 @@ plot_tool=SimpleMatplotlibTool() # Matplotlib Tool which will be used by plotter worker ) -agent_server = AgentServer(agent=baseline, port={port}) +agent_server = AgentServer( + agent=baseline, + port={port}, + api_token=os.environ.get("PREMSQL_API_TOKEN") +) agent_server.launch() """ STARTER_CODE_FILE_MLX = """ +import os from premsql.playground import AgentServer from premsql.agents import BaseLineAgent from premsql.generators import Text2SQLGeneratorMLX @@ -42,6 +47,7 @@ STARTER_CODE_FILE_OLLAMA = """ +import os from premsql.playground import AgentServer from premsql.agents import BaseLineAgent from premsql.generators import Text2SQLGeneratorOllama @@ -63,6 +69,7 @@ STARTER_CODE_FILE_HF = """ +import os from premsql.playground import AgentServer from premsql.agents import BaseLineAgent from premsql.generators import Text2SQLGeneratorHF @@ -192,9 +199,11 @@ def render_kaggle_view() -> Tuple[Optional[str], Optional[Path]]: if submit: if not session_name: st.error("Please enter a session name") - + return None, None + if not _is_valid_kaggle_id(kaggle_id): st.error("Invalid Kaggle Id") + return None, None try: with st.spinner(text="Downloading from Kaggle"): @@ -229,9 +238,11 @@ def render_csv_upload_view() -> Tuple[Optional[str], Optional[Path]]: if submit: if not session_name: st.error("Please enter a session name") - + return None, None + if not uploaded_files: st.error("Please upload at least one CSV file") + return None, None try: with st.spinner(text="Processing CSV files"): diff --git a/premsql/playground/frontend/main.py b/premsql/playground/frontend/main.py index b9cef9b..bfe967a 100644 --- a/premsql/playground/frontend/main.py +++ b/premsql/playground/frontend/main.py @@ -1,53 +1,68 @@ +import os +from pathlib import Path + +# Load environment variables from .env file before importing other modules +try: + from dotenv import load_dotenv + env_file = Path(__file__).parent.parent.parent.parent / ".env" + if env_file.exists(): + load_dotenv(env_file) +except ImportError: + pass + import streamlit as st -from premsql.playground.frontend.components.chat import ChatComponent +from premsql.playground.frontend.components.chat import ChatComponent from premsql.playground.frontend.components.session import SessionComponent from premsql.playground.frontend.components.uploader import UploadComponent st.set_page_config(page_title="PremSQL Playground", page_icon="🔍", layout="wide") + def render_main_view(): session_component = SessionComponent() - - selected_session = session_component.render_list_sessions() + + # Render session list with integrated delete buttons + selected_session = session_component.render_session_list_with_actions() + + # Render create new session form session_creation = session_component.render_register_session() + + # Render help links session_component.render_additional_links() - if session_creation is not None: - if session_creation.status_code == 200: - new_session_name = session_creation.session_name - st.success(f"New session created: {new_session_name}") - ChatComponent().render_chat_env(session_name=new_session_name) + # Render chat for selected or newly created session + if session_creation is not None and session_creation.status_code == 200: + ChatComponent().render_chat_env(session_name=session_creation.session_name) elif selected_session is not None: ChatComponent().render_chat_env(session_name=selected_session) - - session_component.render_delete_session_view() + def main(): + # Header with logo _, col2, _ = st.sidebar.columns([1, 2, 1]) with col2: st.image( "https://static.premai.io/logo.svg", use_container_width=True, - width=150, - clamp=True, ) st.header("PremSQL Playground") + st.title("PremSQL Playground") - - # Add navigation - selected_page = st.sidebar.selectbox("Navigation", ["Chat", "Upload csvs or use Kaggle"]) - + + # Navigation + selected_page = st.sidebar.selectbox("Navigation", ["Chat", "Upload Data"]) + if selected_page == "Chat": - st.write("Welcome to the PremSQL Playground. Select or create a session to get started.") + st.write("Select or create a session to start chatting with your database.") render_main_view() else: st.write( - "You can either upload multiple csv files or enter a valid Kaggle ID. " - "This will migrate all the csvs into a SQLite Database. You can then " - "use them for natural language powered analysis using PremSQL." + "Upload CSV files or enter a Kaggle dataset ID to create a SQLite database " + "for natural language powered analysis." ) UploadComponent.render_kaggle_view() UploadComponent.render_csv_upload_view() + if __name__ == "__main__": main() \ No newline at end of file diff --git a/premsql/playground/frontend/utils.py b/premsql/playground/frontend/utils.py index 86eaa87..e4c2622 100644 --- a/premsql/playground/frontend/utils.py +++ b/premsql/playground/frontend/utils.py @@ -6,9 +6,21 @@ from pathlib import Path from platformdirs import user_cache_dir from premsql.logger import setup_console_logger +from premsql.security import ( + resolve_path_within_root, + sanitize_filename, + validate_session_name, +) logger = setup_console_logger("[FRONTEND-UTILS]") +MAX_UPLOAD_FILES = 20 +MAX_FILE_SIZE_BYTES = 25 * 1024 * 1024 +MAX_TOTAL_UPLOAD_SIZE_BYTES = 100 * 1024 * 1024 +MAX_CSV_ROWS = 250_000 +MAX_CSV_COLUMNS = 200 +MAX_CSV_FILES_PER_IMPORT = 50 + def _is_valid_kaggle_id(kaggle_id: str) -> bool: pattern = r'^[a-zA-Z0-9_-]+/[a-zA-Z0-9_-]+$' return bool(re.match(pattern, kaggle_id)) @@ -21,9 +33,27 @@ def _migrate_to_sqlite(csv_folder: Path, sqlite_db_path: Path) -> Path: """Common migration logic for both Kaggle and local CSV uploads.""" conn = sqlite3.connect(sqlite_db_path) try: - for csv_file in csv_folder.glob('*.csv'): + csv_files = list(csv_folder.glob("*.csv")) + if len(csv_files) > MAX_CSV_FILES_PER_IMPORT: + raise ValueError( + f"Too many CSV files. Maximum allowed is {MAX_CSV_FILES_PER_IMPORT}." + ) + + for csv_file in csv_files: + if csv_file.stat().st_size > MAX_FILE_SIZE_BYTES: + raise ValueError( + f"File '{csv_file.name}' exceeds the {MAX_FILE_SIZE_BYTES // (1024 * 1024)}MB limit." + ) table_name = csv_file.stem df = pd.read_csv(csv_file) + if len(df) > MAX_CSV_ROWS: + raise ValueError( + f"File '{csv_file.name}' exceeds the maximum row limit of {MAX_CSV_ROWS}." + ) + if len(df.columns) > MAX_CSV_COLUMNS: + raise ValueError( + f"File '{csv_file.name}' exceeds the maximum column limit of {MAX_CSV_COLUMNS}." + ) df.to_sql(table_name, conn, if_exists='replace', index=False) logger.info(f"Migrated {csv_file.name} to table '{table_name}'") @@ -39,28 +69,51 @@ def migrate_from_csv_to_sqlite( folder_containing_csvs: str, session_name: str ) -> Path: + validated_session_name = validate_session_name(session_name) sqlite_db_folder = Path(user_cache_dir()) / "premsql" / "kaggle" os.makedirs(sqlite_db_folder, exist_ok=True) - sqlite_db_path = sqlite_db_folder / f"{session_name}.sqlite" + sqlite_db_path = resolve_path_within_root( + sqlite_db_folder, f"{validated_session_name}.sqlite" + ) return _migrate_to_sqlite(Path(folder_containing_csvs), sqlite_db_path) def migrate_local_csvs_to_sqlite( uploaded_files: list, session_name: str ) -> Path: + validated_session_name = validate_session_name(session_name) cache_dir = Path(user_cache_dir()) - csv_folder = cache_dir / "premsql" / "csv_uploads" / session_name - sqlite_db_folder = cache_dir / "premsql" / "csv_uploads" + csv_root = cache_dir / "premsql" / "csv_uploads" + csv_folder = resolve_path_within_root(csv_root, validated_session_name) + sqlite_db_folder = csv_root os.makedirs(csv_folder, exist_ok=True) os.makedirs(sqlite_db_folder, exist_ok=True) - sqlite_db_path = sqlite_db_folder / f"{session_name}.sqlite" + sqlite_db_path = resolve_path_within_root( + sqlite_db_folder, f"{validated_session_name}.sqlite" + ) + + if len(uploaded_files) > MAX_UPLOAD_FILES: + raise ValueError(f"Too many files uploaded. Maximum allowed is {MAX_UPLOAD_FILES}.") + + total_size = 0 # Save uploaded files to CSV folder for uploaded_file in uploaded_files: - file_path = csv_folder / uploaded_file.name + file_size = getattr(uploaded_file, "size", len(uploaded_file.getvalue())) + if file_size > MAX_FILE_SIZE_BYTES: + raise ValueError( + f"File '{uploaded_file.name}' exceeds the {MAX_FILE_SIZE_BYTES // (1024 * 1024)}MB limit." + ) + total_size += file_size + if total_size > MAX_TOTAL_UPLOAD_SIZE_BYTES: + raise ValueError( + f"Total upload size exceeds the {MAX_TOTAL_UPLOAD_SIZE_BYTES // (1024 * 1024)}MB limit." + ) + safe_name = sanitize_filename(uploaded_file.name) + file_path = resolve_path_within_root(csv_folder, safe_name) with open(file_path, 'wb') as f: f.write(uploaded_file.getvalue()) - return _migrate_to_sqlite(csv_folder, sqlite_db_path) \ No newline at end of file + return _migrate_to_sqlite(csv_folder, sqlite_db_path) diff --git a/premsql/playground/inference_server/api_client.py b/premsql/playground/inference_server/api_client.py index 1572768..2574938 100644 --- a/premsql/playground/inference_server/api_client.py +++ b/premsql/playground/inference_server/api_client.py @@ -3,17 +3,22 @@ import requests +from premsql.security import build_auth_headers + class InferenceServerAPIError(Exception): pass class InferenceServerAPIClient: - def __init__(self, timeout: int = 600) -> None: - self.headers = { + def __init__(self, timeout: int = 120, api_token: Optional[str] = None) -> None: + self.headers = build_auth_headers( + { "accept": "application/json", "Content-Type": "application/json", - } + }, + token=api_token, + ) self.timeout = timeout def _make_request( @@ -59,5 +64,5 @@ def get_chat_history(self, base_url: str, message_id: int) -> Dict[str, Any]: return self._make_request(base_url, "GET", endpoint) def delete_session(self, base_url: str) -> Dict[str, Any]: - endpoint = "/delete_session/" + endpoint = "/session" return self._make_request(base_url, "DELETE", endpoint) diff --git a/premsql/playground/inference_server/service.py b/premsql/playground/inference_server/service.py index 02d9171..6176e29 100644 --- a/premsql/playground/inference_server/service.py +++ b/premsql/playground/inference_server/service.py @@ -1,21 +1,31 @@ -import traceback - from contextlib import asynccontextmanager from datetime import datetime +import os from typing import Optional -from fastapi import FastAPI, HTTPException +from fastapi import FastAPI, HTTPException, Request from fastapi.middleware.cors import CORSMiddleware -from pydantic import BaseModel +from pydantic import BaseModel, Field from premsql.logger import setup_console_logger from premsql.agents.base import AgentBase, AgentOutput +from premsql.security import ( + PREMSQL_API_TOKEN_HEADER, + PremSQLException, + get_allowed_origins, + get_api_token, + redact_agent_output_payload, + safe_error_message, +) logger = setup_console_logger("[FASTAPI-INFERENCE-SERVICE]") +# Check if in debug mode +DJANGO_DEBUG = os.environ.get("PREMSQL_DJANGO_DEBUG", "false").lower() == "true" + class QuestionInput(BaseModel): - question: str + question: str = Field(..., min_length=1, max_length=4000) class SessionInfoResponse(BaseModel): @@ -43,12 +53,31 @@ def __init__( agent: AgentBase, url: Optional[str] = "localhost", port: Optional[int] = 8100, + api_token: Optional[str] = None, ) -> None: self.agent = agent self.port = port self.url = url + # Always get a token (either from config or auto-generated dev token) + self.api_token = get_api_token(api_token) + # Track if this is an auto-generated dev token + self._is_dev_token = api_token is None and not self._check_env_token() self.app = self.create_app() + def _check_env_token(self) -> bool: + """Check if token was set in environment""" + import os + return os.environ.get("PREMSQL_API_TOKEN") is not None + + def _authorize_request(self, request: Request) -> None: + # In debug mode, skip authentication + if DJANGO_DEBUG: + return + if not self.api_token: + return + if request.headers.get(PREMSQL_API_TOKEN_HEADER) != self.api_token: + raise HTTPException(status_code=401, detail="Unauthorized") + @asynccontextmanager async def lifespan(self, app: FastAPI): # Startup: Log the initialization @@ -63,32 +92,36 @@ def create_app(self): app = FastAPI(lifespan=self.lifespan) app.add_middleware( CORSMiddleware, - allow_origins=["*"], # Allows all origins - allow_credentials=True, - allow_methods=["*"], # Allows all methods - allow_headers=["*"], # Allows all headers + allow_origins=get_allowed_origins(), + allow_credentials=False, + allow_methods=["GET", "POST", "DELETE"], + allow_headers=["Content-Type", PREMSQL_API_TOKEN_HEADER], ) @app.post("/completion", response_model=CompletionResponse) - async def completion(input_data: QuestionInput): + async def completion(input_data: QuestionInput, request: Request): + self._authorize_request(request) try: result = self.agent(question=input_data.question, server_mode=True) message_id = self.agent.history.get_latest_message_id() + payload = redact_agent_output_payload(result.model_dump()) return CompletionResponse( - message=AgentOutput(**result.model_dump()), message_id=message_id + message=AgentOutput(**payload), message_id=message_id ) - except Exception as e: - stack_trace = traceback.format_exc() - logger.error(stack_trace) - logger.error(f"Error processing query: {str(e)}") + except Exception as exc: + # Use safe error message for logging (whitelist approach) + logger.error(safe_error_message(exc, debug_mode=False)) + # Return safe message to user raise HTTPException( - status_code=500, detail=f"Error processing query: {str(e)}" + status_code=500, + detail=safe_error_message(exc, debug_mode=False) ) # TODO: I need a method which will just get the "latets message_id" @app.get("/chat_history/{message_id}", response_model=ChatHistoryResponse) - async def get_chat_history(message_id: int): + async def get_chat_history(message_id: int, request: Request): + self._authorize_request(request) try: exit_output = self.agent.history.get_by_message_id( message_id=message_id @@ -101,49 +134,48 @@ async def get_chat_history(message_id: int): agent_output = self.agent.convert_exit_output_to_agent_output( exit_output=exit_output ) + payload = redact_agent_output_payload(agent_output.model_dump()) return ChatHistoryResponse( - message_id=message_id, agent_output=agent_output + message_id=message_id, agent_output=AgentOutput(**payload) ) - except Exception as e: - logger.error(f"Error retrieving chat history: {str(e)}") + except Exception as exc: + logger.error(safe_error_message(exc, debug_mode=False)) raise HTTPException( - status_code=500, detail=f"Error retrieving chat history: {str(e)}" + status_code=500, + detail=safe_error_message(exc, debug_mode=False) ) @app.get("/") - async def health_check(): + async def health_check(request: Request): + self._authorize_request(request) return { "status_code": 200, "status": f"healthy, running: {self.agent.session_name}" } @app.get("/health") - async def health_check(): + async def health_check_no_auth(request: Request): + # Health check doesn't require authentication return {"status_code": 200, "status": "healthy"} @app.get("/session_info", response_model=SessionInfoResponse) - async def get_session_info(): + async def get_session_info(request: Request): + self._authorize_request(request) try: session_name = getattr(self.agent, "session_name", None) - db_connection_uri = getattr(self.agent, "db_connection_uri", None) - session_db_path = getattr(self.agent, "session_db_path", None) - - if any( - attr is None - for attr in [session_name, db_connection_uri, session_db_path] - ): - raise ValueError("One or more required attributes are None") + if session_name is None: + raise ValueError("Session name is not available") return SessionInfoResponse( status=200, session_name=session_name, - db_connection_uri=db_connection_uri, - session_db_path=session_db_path, - base_url=f"{self.url}:{self.port}", + db_connection_uri=None, + session_db_path=None, + base_url=f"http://{self.url}:{self.port}", created_at=datetime.now(), ) - except Exception as e: - logger.error(f"Error getting session info: {str(e)}") + except Exception as exc: + logger.error(safe_error_message(exc, debug_mode=False)) return SessionInfoResponse( status=500, session_name=None, @@ -153,10 +185,29 @@ async def get_session_info(): created_at=None, ) + @app.delete("/session") + async def delete_session(request: Request): + self._authorize_request(request) + try: + self.agent.history.delete_table() + # Recreate the table for future queries + self.agent.history.create_table_if_not_exists() + return {"status_code": 200, "status": "success"} + except Exception as exc: + logger.error(safe_error_message(exc, debug_mode=False)) + raise HTTPException( + status_code=500, + detail=safe_error_message(exc, debug_mode=False) + ) + return app def launch(self): import uvicorn logger.info(f"Starting server on port {self.port}") + if self._is_dev_token: + logger.info(f"Using auto-generated dev token: {self.api_token}") + else: + logger.info(f"Using configured API token") uvicorn.run(self.app, host=self.url, port=int(self.port)) diff --git a/premsql/prompts.py b/premsql/prompts.py index 1ef083d..c278290 100644 --- a/premsql/prompts.py +++ b/premsql/prompts.py @@ -6,7 +6,8 @@ 1. Do not add ``` at start / end of the query. It should be a single line query in a single line (string format) 2. Make sure the column names are correct and exists in the table 3. For column names which has a space with it, make sure you have put `` in that column name -4. Think step by step and always check schema and question and the column names before writing the +4. Only generate read-only SQL queries. Never generate INSERT, UPDATE, DELETE, DROP, ALTER, or PRAGMA statements. +5. Think step by step and always check schema and question and the column names before writing the query. # Database and Table Schema: @@ -30,6 +31,7 @@ a single line (string format). - Make sure the column names are correct and exists in the table - For column names which has a space with it, make sure you have put `` in that column name +- Only generate read-only SQL queries. Never generate INSERT, UPDATE, DELETE, DROP, ALTER, or PRAGMA statements. # Database and Table Schema: {schemas} diff --git a/premsql/security.py b/premsql/security.py new file mode 100644 index 0000000..2ce4b2c --- /dev/null +++ b/premsql/security.py @@ -0,0 +1,319 @@ +import ast +import json +import logging +import os +import re +import secrets +from pathlib import Path +from typing import Any, Iterable, Optional +from urllib.parse import urlparse + +try: + import sqlparse +except ImportError: # pragma: no cover - fallback exercised in minimal envs + sqlparse = None + +SESSION_NAME_PATTERN = re.compile(r"^[A-Za-z0-9_-]{1,64}$") +SAFE_FILENAME_PATTERN = re.compile(r"[^A-Za-z0-9._-]+") +FORBIDDEN_SQL_PATTERN = re.compile( + r"\b(" + r"INSERT|UPDATE|DELETE|DROP|ALTER|ATTACH|DETACH|PRAGMA|REPLACE|" + r"CREATE|TRUNCATE|MERGE|GRANT|REVOKE|VACUUM|ANALYZE|COPY|CALL" + r")\b", + re.IGNORECASE, +) +CODE_FENCE_PATTERN = re.compile(r"^```(?:json|python)?\s*(.*?)\s*```$", re.DOTALL) +LOOPBACK_HOSTS = {"127.0.0.1", "localhost", "::1"} +DEFAULT_ALLOWED_ORIGINS = ( + "http://127.0.0.1:8501", + "http://localhost:8501", +) +PREMSQL_API_TOKEN_HEADER = "X-PremSQL-API-Token" +PREMSQL_API_TOKEN_ENV = "PREMSQL_API_TOKEN" + +# Auto-generated dev token (only used when PREMSQL_API_TOKEN not configured) +_DEV_API_TOKEN: Optional[str] = None + + +class SecurityValidationError(ValueError): + """Base exception for security validation failures with safe messages.""" + def __init__(self, message: str, *args): + # All security validation errors are safe to expose (they're validation messages) + self.safe_message = message + super().__init__(message, *args) + + +class UnsafeSQLQuery(SecurityValidationError): + """Exception for unsafe SQL queries.""" + pass + + +class PremSQLException(Exception): + """ + Custom exception base class that forces developers to explicitly declare safe messages. + All PremSQL exceptions should inherit from this class for production use. + """ + def __init__(self, safe_message: str, internal_message: Optional[str] = None, *args): + self.safe_message = safe_message # Safe message to expose to users/logs + self.internal_message = internal_message or safe_message # For internal debugging + super().__init__(internal_message or safe_message, *args) + + +def safe_error_message(exc: Exception, debug_mode: bool = False) -> str: + """ + Generate a safe error message for external exposure using pure whitelist approach. + + Security principle: Never try to filter/strip bad content (blacklist). + Only expose what we explicitly know is safe (whitelist). + + Args: + exc: The exception to generate message from + debug_mode: If True, include more details (only for trusted environments) + + Returns: + A safe message that doesn't expose sensitive information + """ + # Whitelist approach: only expose what we explicitly allow + + # 1. PremSQLException and SecurityValidationError have explicitly declared safe messages + if isinstance(exc, (PremSQLException, SecurityValidationError)): + return exc.safe_message + + # 2. HTTPException from FastAPI - use the detail field which is designed for user display + if hasattr(exc, 'detail') and isinstance(exc.detail, str): + return exc.detail + + # 3. In debug mode (trusted environment only), allow more details + if debug_mode: + # Still only for known safe exception types + if isinstance(exc, (ValueError, TypeError, KeyError)): + msg = str(exc) + return msg[:200] if len(msg) > 200 else msg + + # 4. For all other exceptions: DO NOT expose any message content + # Only expose the exception type name - this is safe metadata + return f"An error occurred ({type(exc).__name__}). Please try again or contact support." + + +def validate_session_name(session_name: str) -> str: + if session_name is None: + raise SecurityValidationError("Session name is required") + + normalized = session_name.strip() + if not SESSION_NAME_PATTERN.fullmatch(normalized): + raise SecurityValidationError( + "Session name must contain only letters, numbers, underscores, or hyphens " + "and be at most 64 characters long" + ) + return normalized + + +def quote_sqlite_identifier(identifier: str) -> str: + return f'"{validate_session_name(identifier)}"' + + +def sanitize_filename(filename: str) -> str: + if filename is None: + raise SecurityValidationError("Filename is required") + + basename = Path(filename).name.strip() + if not basename or basename in {".", ".."}: + raise SecurityValidationError("Filename is invalid") + + sanitized = SAFE_FILENAME_PATTERN.sub("_", basename) + if not sanitized or sanitized in {".", ".."}: + raise SecurityValidationError("Filename is invalid after sanitization") + return sanitized + + +def resolve_path_within_root(root: Path | str, *parts: str) -> Path: + root_path = Path(root).resolve() + target = root_path.joinpath(*parts).resolve() + if target != root_path and root_path not in target.parents: + raise SecurityValidationError("Resolved path escapes the configured root") + return target + + +def normalize_base_url(base_url: str) -> str: + if not base_url: + raise SecurityValidationError("Base URL is required") + + candidate = base_url.strip() + if "://" not in candidate: + candidate = f"http://{candidate}" + + parsed = urlparse(candidate) + if parsed.scheme not in {"http", "https"}: + raise SecurityValidationError("Base URL must use http or https") + if parsed.hostname not in LOOPBACK_HOSTS: + raise SecurityValidationError( + "Base URL must point to a local loopback address" + ) + if parsed.username or parsed.password: + raise SecurityValidationError("Credentials are not allowed in base URLs") + if parsed.path not in {"", "/"} or parsed.params or parsed.query or parsed.fragment: + raise SecurityValidationError( + "Base URL must only include scheme, host, and port" + ) + if parsed.port is None: + raise SecurityValidationError("Base URL must include an explicit port") + + host = parsed.hostname + if ":" in host and not host.startswith("["): + host = f"[{host}]" + return f"{parsed.scheme}://{host}:{parsed.port}" + + +def strip_code_fences(content: str) -> str: + if content is None: + raise SecurityValidationError("Model output is empty") + + cleaned = content.strip() + match = CODE_FENCE_PATTERN.match(cleaned) + if match: + return match.group(1).strip() + return cleaned + + +def parse_structured_output( + content: str, expected_keys: Optional[Iterable[str]] = None +) -> dict[str, Any]: + cleaned = strip_code_fences(content) + + try: + parsed = json.loads(cleaned) + except json.JSONDecodeError: + pythonish = ( + cleaned.replace("null", "None") + .replace("true", "True") + .replace("false", "False") + ) + try: + parsed = ast.literal_eval(pythonish) + except (ValueError, SyntaxError) as exc: + raise SecurityValidationError( + "Model output must be valid JSON" + ) from exc + + if not isinstance(parsed, dict): + raise SecurityValidationError("Structured model output must be a JSON object") + + if expected_keys is not None: + expected = set(expected_keys) + missing = expected - set(parsed) + if missing: + raise SecurityValidationError( + f"Structured model output is missing keys: {', '.join(sorted(missing))}" + ) + return parsed + + +def ensure_expected_keys_only( + data: dict[str, Any], expected_keys: Iterable[str] +) -> dict[str, Any]: + expected = set(expected_keys) + unexpected = set(data) - expected + if unexpected: + raise SecurityValidationError( + f"Structured model output contains unexpected keys: " + f"{', '.join(sorted(unexpected))}" + ) + return data + + +def normalize_sql(sql: str) -> str: + if not sql or not sql.strip(): + raise UnsafeSQLQuery("SQL query is empty") + + if sqlparse is not None: + cleaned = sqlparse.format(sql, strip_comments=True).strip() + statements = [ + statement.strip() for statement in sqlparse.split(cleaned) if statement.strip() + ] + else: + cleaned = re.sub(r"--.*?$", "", sql, flags=re.MULTILINE).strip() + statements = [statement.strip() for statement in cleaned.split(";") if statement.strip()] + if len(statements) != 1: + raise UnsafeSQLQuery("Only a single SQL statement is allowed") + return statements[0] + + +def enforce_read_only_sql(sql: str) -> str: + normalized = normalize_sql(sql) + upper_sql = normalized.upper() + if not (upper_sql.startswith("SELECT") or upper_sql.startswith("WITH")): + raise UnsafeSQLQuery("Only read-only SELECT queries are allowed") + if FORBIDDEN_SQL_PATTERN.search(upper_sql): + raise UnsafeSQLQuery("Potentially destructive SQL statements are not allowed") + return normalized + + +def mask_db_connection_uri(db_connection_uri: Optional[str]) -> Optional[str]: + if not db_connection_uri: + return None + + if db_connection_uri.startswith("sqlite:///"): + return "sqlite:///***" + + parsed = urlparse(db_connection_uri) + if parsed.scheme: + return f"{parsed.scheme}://***" + return "***" + + +def redact_agent_output_payload(payload: Optional[dict[str, Any]]) -> Optional[dict[str, Any]]: + if payload is None: + return None + + sanitized = dict(payload) + sanitized["db_connection_uri"] = None + return sanitized + + +def get_api_token(explicit_token: Optional[str] = None) -> Optional[str]: + """ + Get API token from explicit parameter, environment, or auto-generate for dev. + + For production (DEBUG=false): PREMSQL_API_TOKEN must be configured + For development (DEBUG=true): Auto-generates a random token if not configured + """ + token = explicit_token or os.environ.get(PREMSQL_API_TOKEN_ENV) + if token: + return token + + # Check if in debug mode + debug_mode = os.environ.get("PREMSQL_DJANGO_DEBUG", "false").lower() == "true" + + if not debug_mode: + # Production mode: Token must be configured + raise ValueError( + "PREMSQL_API_TOKEN must be configured for production. " + "Set PREMSQL_API_TOKEN environment variable or set PREMSQL_DJANGO_DEBUG=true for development." + ) + + # Development mode: Auto-generate dev token + global _DEV_API_TOKEN + if _DEV_API_TOKEN is None: + _DEV_API_TOKEN = f"dev-{secrets.token_hex(16)}" + logging.getLogger("premsql.security").warning( + f"Auto-generated dev API token: {_DEV_API_TOKEN} " + "(Set PREMSQL_API_TOKEN for production)" + ) + return _DEV_API_TOKEN + + +def build_auth_headers( + headers: Optional[dict[str, str]] = None, token: Optional[str] = None +) -> dict[str, str]: + merged = dict(headers or {}) + api_token = get_api_token(token) + if api_token: + merged[PREMSQL_API_TOKEN_HEADER] = api_token + return merged + + +def get_allowed_origins() -> list[str]: + configured = os.environ.get("PREMSQL_ALLOWED_ORIGINS") + if not configured: + return list(DEFAULT_ALLOWED_ORIGINS) + return [origin.strip() for origin in configured.split(",") if origin.strip()] diff --git a/premsql/utils.py b/premsql/utils.py index 7f360e4..9947e9f 100644 --- a/premsql/utils.py +++ b/premsql/utils.py @@ -78,8 +78,10 @@ def sqlite_schema_prompt(db_path: str) -> str: table_name = table[0] if table_name == "sqlite_sequence": continue + # Use parameterized query instead of f-string for safety cursor.execute( - f"SELECT sql FROM sqlite_master WHERE type='table' AND name='{table_name}';" + "SELECT sql FROM sqlite_master WHERE type='table' AND name=?;", + (table_name,) ) create_table_sql = cursor.fetchone() if create_table_sql: diff --git a/pyproject.toml b/pyproject.toml index abe5dbb..aceed3a 100644 --- a/pyproject.toml +++ b/pyproject.toml @@ -6,11 +6,11 @@ authors = ["Anindyadeep "] readme = "README.md" [tool.poetry.dependencies] -python = "^3.10" +python = ">=3.10,<3.13" datasets = "^2.20.0" einops = "^0.8.0" black = "^24.4.2" -fastapi = "^0.112.0" +fastapi = ">=0.115.0" huggingface-hub = "^0.24.5" isort = "^5.13.2" numpy = "^1.26.3" @@ -28,9 +28,11 @@ drf-yasg = "^1.21.8" func_timeout = "^4.3.5" matplotlib = "^3.9.2" pillow = ">=8,<11" -uvicorn = "^0.32.0" +uvicorn = ">=0.32.0" streamlit = "^1.40.0" kagglehub = "^0.3.3" +httpx = ">=0.27.0" +starlette = ">=0.41.0" [tool.poetry.extras] mac = ["mlx", "mlx-lm"] diff --git a/start_agent.py b/start_agent.py new file mode 100644 index 0000000..d49c2c8 --- /dev/null +++ b/start_agent.py @@ -0,0 +1,332 @@ +""" +PremSQL Agent Server Startup Script + +This script starts the AgentServer which handles text-to-SQL queries. +You need to configure an LLM provider (OpenAI, PremAI, vLLM, or Ollama). + +Usage: +1. Set your API key in .env file +2. Run: python start_agent.py + +Supported LLM Providers: +- vLLM (self-hosted models via OpenAI-compatible API) +- OpenAI (GPT-4, GPT-3.5, etc.) +- PremAI (Prem's hosted models) +- Ollama (local models) +- Custom (any OpenAI-compatible service) +""" + +import os +import sys +import uuid +from pathlib import Path +from dotenv import load_dotenv + +# Load environment variables +load_dotenv() + +# Add project to path +project_root = Path(__file__).parent +sys.path.insert(0, str(project_root)) + +from premsql.playground import AgentServer +from premsql.agents import BaseLineAgent +from premsql.executors import ExecutorUsingLangChain +from premsql.agents.tools import SimpleMatplotlibTool +from premsql.security import get_api_token + +# Configuration +SESSION_NAME = os.environ.get("PREMSQL_SESSION_NAME", f"session_{uuid.uuid4().hex[:8]}") +DB_PATH = project_root / "sample_data" / "schools.db" +DB_CONNECTION_URI = f"sqlite:///{DB_PATH}" +PORT = int(os.environ.get("PREMSQL_AGENT_PORT", 8100)) +# Get actual token (will auto-generate if not configured) +actual_token = get_api_token() + +print(f"Database: {DB_CONNECTION_URI}") +print(f"Session: {SESSION_NAME}") +print(f"Port: {PORT}") + + +def create_vllm_agent(): + """Create agent using vLLM deployed model""" + from premsql.generators import Text2SQLGeneratorVLLM + + # Support both old and new config names + base_url = os.environ.get("VLLM_BASE_URL") or os.environ.get("VLLM_ENDPOINT") + model_name = os.environ.get("VLLM_MODEL_NAME", "default") + + if not base_url: + raise ValueError("VLLM_BASE_URL is required. Set it in .env file.") + + print(f"Using vLLM at: {base_url}") + print(f"Model: {model_name}") + + text2sql_model = Text2SQLGeneratorVLLM( + model_name=model_name, + experiment_name="text2sql_vllm", + type="test", + base_url=base_url, + # disable_thinking auto-detected for Qwen models + ) + + analyser_model = Text2SQLGeneratorVLLM( + model_name=model_name, + experiment_name="analyser_vllm", + type="test", + base_url=base_url, + ) + + return BaseLineAgent( + session_name=SESSION_NAME, + db_connection_uri=DB_CONNECTION_URI, + specialized_model1=text2sql_model, + specialized_model2=analyser_model, + executor=ExecutorUsingLangChain(), + auto_filter_tables=False, + plot_tool=SimpleMatplotlibTool() + ) + + +def create_custom_agent(): + """Create agent using any OpenAI-compatible service""" + from premsql.generators import Text2SQLGeneratorOpenAICompatible + + base_url = os.environ.get("CUSTOM_BASE_URL") + model_name = os.environ.get("CUSTOM_MODEL_NAME", "default") + api_key = os.environ.get("CUSTOM_API_KEY", "dummy-key") + + if not base_url: + raise ValueError("CUSTOM_BASE_URL is required for custom provider.") + + print(f"Using custom service at: {base_url}") + print(f"Model: {model_name}") + + # Optional: extra_body for service-specific parameters + extra_body = None + if os.environ.get("CUSTOM_DISABLE_THINKING", "").lower() == "true": + extra_body = {"chat_template_kwargs": {"enable_thinking": False}} + + text2sql_model = Text2SQLGeneratorOpenAICompatible( + model_name=model_name, + experiment_name="text2sql_custom", + type="test", + base_url=base_url, + api_key=api_key, + extra_body=extra_body, + ) + + analyser_model = Text2SQLGeneratorOpenAICompatible( + model_name=model_name, + experiment_name="analyser_custom", + type="test", + base_url=base_url, + api_key=api_key, + extra_body=extra_body, + ) + + return BaseLineAgent( + session_name=SESSION_NAME, + db_connection_uri=DB_CONNECTION_URI, + specialized_model1=text2sql_model, + specialized_model2=analyser_model, + executor=ExecutorUsingLangChain(), + auto_filter_tables=False, + plot_tool=SimpleMatplotlibTool() + ) + + +def create_openai_agent(): + """Create agent using OpenAI GPT model""" + from premsql.generators import Text2SQLGeneratorOpenAI + + api_key = os.environ.get("OPENAI_API_KEY") + if not api_key or api_key == "your-openai-api-key-here": + raise ValueError("Please set OPENAI_API_KEY in .env file") + + model_name = os.environ.get("OPENAI_MODEL_NAME", "gpt-4o-mini") + print(f"Using OpenAI model: {model_name}") + + text2sql_model = Text2SQLGeneratorOpenAI( + model_name=model_name, + experiment_name="text2sql_openai", + type="test", + openai_api_key=api_key + ) + + analyser_model = Text2SQLGeneratorOpenAI( + model_name=model_name, + experiment_name="analyser_openai", + type="test", + openai_api_key=api_key + ) + + return BaseLineAgent( + session_name=SESSION_NAME, + db_connection_uri=DB_CONNECTION_URI, + specialized_model1=text2sql_model, + specialized_model2=analyser_model, + executor=ExecutorUsingLangChain(), + auto_filter_tables=False, + plot_tool=SimpleMatplotlibTool() + ) + + +def create_premai_agent(): + """Create agent using PremAI model""" + from premsql.generators import Text2SQLGeneratorPremAI + + api_key = os.environ.get("PREMAI_API_KEY") + project_id = os.environ.get("PREMAI_PROJECT_ID") + + if not api_key or api_key == "your-premai-api-key-here": + raise ValueError("Please set PREMAI_API_KEY in .env file") + + model_name = os.environ.get("PREMAI_MODEL_NAME", "gpt-4o") + print(f"Using PremAI model: {model_name}") + + text2sql_model = Text2SQLGeneratorPremAI( + model_name=model_name, + experiment_name="text2sql_premai", + type="test", + premai_api_key=api_key, + project_id=project_id + ) + + analyser_model = Text2SQLGeneratorPremAI( + model_name=model_name, + experiment_name="analyser_premai", + type="test", + premai_api_key=api_key, + project_id=project_id + ) + + return BaseLineAgent( + session_name=SESSION_NAME, + db_connection_uri=DB_CONNECTION_URI, + specialized_model1=text2sql_model, + specialized_model2=analyser_model, + executor=ExecutorUsingLangChain(), + auto_filter_tables=False, + plot_tool=SimpleMatplotlibTool() + ) + + +def create_ollama_agent(): + """Create agent using Ollama local model""" + from premsql.generators import Text2SQLGeneratorOllama + + base_url = os.environ.get("OLLAMA_BASE_URL", "http://127.0.0.1:11434") + model_name = os.environ.get("OLLAMA_MODEL_NAME", "llama3.2") + + print(f"Using Ollama at: {base_url}") + print(f"Model: {model_name}") + + text2sql_model = Text2SQLGeneratorOllama( + model_name=model_name, + experiment_name="text2sql_ollama", + type="test", + base_url=base_url + ) + + analyser_model = Text2SQLGeneratorOllama( + model_name=model_name, + experiment_name="analyser_ollama", + type="test", + base_url=base_url + ) + + return BaseLineAgent( + session_name=SESSION_NAME, + db_connection_uri=DB_CONNECTION_URI, + specialized_model1=text2sql_model, + specialized_model2=analyser_model, + executor=ExecutorUsingLangChain(), + auto_filter_tables=False, + plot_tool=SimpleMatplotlibTool() + ) + + +def detect_provider(): + """Auto-detect which LLM provider is configured""" + providers = [] + + # Check vLLM + if os.environ.get("VLLM_BASE_URL"): + providers.append(("vllm", "vLLM")) + + # Check Custom OpenAI-compatible + if os.environ.get("CUSTOM_BASE_URL"): + providers.append(("custom", "Custom OpenAI-compatible")) + + # Check OpenAI + openai_key = os.environ.get("OPENAI_API_KEY", "") + if openai_key and openai_key != "your-openai-api-key-here": + providers.append(("openai", "OpenAI")) + + # Check PremAI + premai_key = os.environ.get("PREMAI_API_KEY", "") + if premai_key and premai_key != "your-premai-api-key-here": + providers.append(("premai", "PremAI")) + + # Check Ollama (always available if installed) + providers.append(("ollama", "Ollama")) + + return providers + + +def main(): + # Detect available providers + providers = detect_provider() + + if len(providers) == 1 and providers[0][0] == "ollama": + # Only Ollama available, check if it's actually running + print("=" * 60) + print("No external LLM provider configured.") + print("=" * 60) + print("\nConfigure one of the following in .env file:") + print(" 1. VLLM_BASE_URL + VLLM_MODEL_NAME - for vLLM deployments") + print(" 2. CUSTOM_BASE_URL + CUSTOM_MODEL_NAME - for any OpenAI-compatible service") + print(" 3. OPENAI_API_KEY + OPENAI_MODEL_NAME - for OpenAI GPT models") + print(" 4. PREMAI_API_KEY + PREMAI_PROJECT_ID - for PremAI") + print("\nOr use Ollama locally: https://ollama.com") + print("=" * 60) + sys.exit(1) + + # Use first configured provider (priority order) + provider_type, provider_name = providers[0] + print(f"Using {provider_name} as LLM provider") + + # Create agent based on provider + creators = { + "vllm": create_vllm_agent, + "custom": create_custom_agent, + "openai": create_openai_agent, + "premai": create_premai_agent, + "ollama": create_ollama_agent, + } + + agent = creators[provider_type]() + + # Start AgentServer + print(f"\nStarting AgentServer on port {PORT}...") + is_dev_token = not os.environ.get("PREMSQL_API_TOKEN") + if is_dev_token: + print(f"API Token (auto-generated): {actual_token}") + else: + print(f"API Token: [configured]") + print(f"\nURL: http://127.0.0.1:{PORT}") + print("=" * 60) + + agent_server = AgentServer( + agent=agent, + url="localhost", + port=PORT, + api_token=None # Let AgentServer handle token internally + ) + + agent_server.launch() + + +if __name__ == "__main__": + main() \ No newline at end of file