""" Unit tests for server.py All tests run without a real Snowflake connection. Snowflake connector and MCP Context are mocked where needed. """ from __future__ import annotations import os import time import unittest from pathlib import Path from unittest.mock import MagicMock, patch import pytest import server from server import ( MAX_OUTPUT_CHARS, MAX_ROWS, MUTATION_TOKEN_TTL, _PENDING_MUTATIONS, _build_connect_kwargs, _get_config, _guard_mutation_sql, _guard_sql, _to_markdown, _validate_context_ids, _validate_identifier, _validate_qualified_identifier, confirm_mutation, describe_table, list_databases, list_schemas, list_tables, preview_mutation, run_query, ) # --------------------------------------------------------------------------- # Helpers # --------------------------------------------------------------------------- def _make_ctx(columns, rows, rowcount=None): """Return a minimal fake MCP Context whose cursor yields the given columns/rows.""" cur = MagicMock() cur.__enter__ = lambda s: s cur.__exit__ = MagicMock(return_value=False) cur.description = [(c,) for c in columns] if columns else None cur.fetchmany.return_value = rows cur.rowcount = rowcount if rowcount is not None else len(rows) conn = MagicMock() conn.cursor.return_value = cur conn.is_closed.return_value = False ctx = MagicMock() ctx.request_context.lifespan_context.conn = conn return ctx def _make_ctx_dml(rowcount=1): """Context that simulates a DML result (no description).""" cur = MagicMock() cur.__enter__ = lambda s: s cur.__exit__ = MagicMock(return_value=False) cur.description = None cur.rowcount = rowcount conn = MagicMock() conn.cursor.return_value = cur conn.is_closed.return_value = False ctx = MagicMock() ctx.request_context.lifespan_context.conn = conn return ctx # --------------------------------------------------------------------------- # _validate_identifier # --------------------------------------------------------------------------- class TestValidateIdentifier(unittest.TestCase): def test_simple_name(self): assert _validate_identifier("ORDERS") == "ORDERS" def test_with_dollar(self): assert _validate_identifier("TBL$1") == "TBL$1" def test_with_underscore(self): assert _validate_identifier("my_table") == "my_table" def test_digits_only(self): assert _validate_identifier("123") == "123" def test_empty_raises(self): with pytest.raises(ValueError, match="Invalid"): _validate_identifier("") def test_dot_raises(self): with pytest.raises(ValueError): _validate_identifier("db.schema") def test_space_raises(self): with pytest.raises(ValueError): _validate_identifier("my table") def test_semicolon_raises(self): with pytest.raises(ValueError): _validate_identifier("t;DROP TABLE x--") def test_single_quote_raises(self): with pytest.raises(ValueError): _validate_identifier("t' OR '1'='1") def test_label_appears_in_message(self): with pytest.raises(ValueError, match="database"): _validate_identifier("bad name", "database") def test_max_length_255_ok(self): name = "A" * 255 assert _validate_identifier(name) == name def test_over_255_raises(self): with pytest.raises(ValueError): _validate_identifier("A" * 256) # --------------------------------------------------------------------------- # _validate_context_ids # --------------------------------------------------------------------------- class TestValidateContextIds(unittest.TestCase): def test_both_none(self): assert _validate_context_ids(None, None) == (None, None) def test_valid_both(self): assert _validate_context_ids("MYDB", "PUBLIC") == ( "MYDB", "PUBLIC" ) def test_database_only(self): db, sc = _validate_context_ids("MYDB", None) assert db == "MYDB" assert sc is None def test_schema_only(self): db, sc = _validate_context_ids(None, "RAW") assert db is None assert sc == "RAW" def test_invalid_database_raises(self): with pytest.raises(ValueError, match="database"): _validate_context_ids("bad db!", None) def test_invalid_schema_raises(self): with pytest.raises(ValueError, match="schema"): _validate_context_ids(None, "bad schema!") # --------------------------------------------------------------------------- # _validate_qualified_identifier # --------------------------------------------------------------------------- class TestValidateQualifiedIdentifier(unittest.TestCase): def test_unqualified(self): assert _validate_qualified_identifier("ORDERS") == "ORDERS" def test_two_part(self): v = "MYDB.PUBLIC" assert _validate_qualified_identifier(v) == v def test_three_part(self): v = "MYDB.PUBLIC.ORDERS" assert _validate_qualified_identifier(v) == v def test_four_parts_raises(self): with pytest.raises(ValueError): _validate_qualified_identifier("A.B.C.D") def test_injection_attempt_raises(self): with pytest.raises(ValueError): _validate_qualified_identifier("x; DROP TABLE y--") def test_space_raises(self): with pytest.raises(ValueError): _validate_qualified_identifier("DB. SCHEMA") # --------------------------------------------------------------------------- # _guard_sql # --------------------------------------------------------------------------- class TestGuardSql(unittest.TestCase): # --- allowed statements --- def test_select(self): stmt = _guard_sql("SELECT 1") assert stmt == "SELECT 1" def test_select_case_insensitive(self): stmt = _guard_sql("select * from t") assert "select" in stmt.lower() def test_show(self): assert _guard_sql("SHOW DATABASES").startswith("SHOW") def test_describe(self): assert _guard_sql("DESCRIBE TABLE t").startswith("DESCRIBE") def test_desc(self): assert _guard_sql("DESC TABLE t").startswith("DESC") def test_explain(self): assert _guard_sql("EXPLAIN SELECT 1").startswith("EXPLAIN") def test_with_cte(self): sql = "WITH cte AS (SELECT 1) SELECT * FROM cte" assert _guard_sql(sql) == sql # --- comment stripping preserves logic --- def test_line_comment_stripped(self): stmt = _guard_sql("SELECT 1 -- drop this comment") assert "--" not in stmt def test_block_comment_stripped(self): stmt = _guard_sql("SELECT /* evil */ 1") assert "/*" not in stmt def test_comment_bypass_rejected(self): # without stripping, "-- \nDROP" would start with "--" # after stripping the comment the remnant is empty with pytest.raises(ValueError): _guard_sql("-- SELECT 1") def test_comment_wrapping_insert_rejected(self): # attacker tries to hide INSERT inside a comment bypass with pytest.raises(ValueError): _guard_sql("/* SELECT */ INSERT INTO t VALUES(1)") # --- blocked statements --- def test_insert_blocked(self): with pytest.raises(ValueError, match="INSERT"): _guard_sql("INSERT INTO t VALUES (1)") def test_update_blocked(self): with pytest.raises(ValueError): _guard_sql("UPDATE t SET x=1") def test_delete_blocked(self): with pytest.raises(ValueError): _guard_sql("DELETE FROM t") def test_drop_blocked(self): with pytest.raises(ValueError): _guard_sql("DROP TABLE t") def test_create_blocked(self): with pytest.raises(ValueError): _guard_sql("CREATE TABLE t (id INT)") def test_grant_blocked(self): with pytest.raises(ValueError): _guard_sql("GRANT SELECT ON t TO ROLE r") # --- multi-statement --- def test_multi_statement_rejected(self): with pytest.raises(ValueError, match="Multi-statement"): _guard_sql("SELECT 1; SELECT 2") def test_empty_raises(self): with pytest.raises(ValueError, match="Empty"): _guard_sql("") def test_only_comment_raises(self): with pytest.raises(ValueError): _guard_sql("/* just a comment */") def test_trailing_semicolon_ok(self): # single statement with trailing ; is fine stmt = _guard_sql("SELECT 1;") assert stmt == "SELECT 1" # --------------------------------------------------------------------------- # _guard_mutation_sql # --------------------------------------------------------------------------- class TestGuardMutationSql(unittest.TestCase): # --- allowed --- def test_insert(self): sql = "INSERT INTO t (a) VALUES (1)" assert _guard_mutation_sql(sql) == sql def test_update(self): sql = "UPDATE t SET a=1 WHERE id=2" assert _guard_mutation_sql(sql) == sql def test_delete(self): sql = "DELETE FROM t WHERE id=1" assert _guard_mutation_sql(sql) == sql def test_merge(self): sql = ( "MERGE INTO t USING s ON t.id=s.id" " WHEN MATCHED THEN UPDATE SET t.v=s.v" ) assert _guard_mutation_sql(sql) == sql def test_call(self): sql = "CALL my_proc(1, 2)" assert _guard_mutation_sql(sql) == sql def test_create_table(self): sql = "CREATE TABLE t (id INT)" assert _guard_mutation_sql(sql) == sql def test_create_or_replace_table(self): sql = "CREATE OR REPLACE TABLE t (id INT)" assert _guard_mutation_sql(sql) == sql def test_create_temp_table(self): sql = "CREATE TEMPORARY TABLE t (id INT)" assert _guard_mutation_sql(sql) == sql def test_create_transient_table(self): sql = "CREATE TRANSIENT TABLE t (id INT)" assert _guard_mutation_sql(sql) == sql # --- blocked --- def test_select_blocked(self): with pytest.raises(ValueError, match="SELECT"): _guard_mutation_sql("SELECT 1") def test_drop_blocked(self): with pytest.raises(ValueError): _guard_mutation_sql("DROP TABLE t") def test_truncate_blocked(self): with pytest.raises(ValueError): _guard_mutation_sql("TRUNCATE TABLE t") def test_alter_blocked(self): with pytest.raises(ValueError): _guard_mutation_sql("ALTER TABLE t ADD COLUMN x INT") def test_create_schema_blocked(self): with pytest.raises(ValueError): _guard_mutation_sql("CREATE SCHEMA s") def test_create_database_blocked(self): with pytest.raises(ValueError): _guard_mutation_sql("CREATE DATABASE d") def test_grant_blocked(self): with pytest.raises(ValueError): _guard_mutation_sql("GRANT SELECT ON t TO ROLE r") def test_revoke_blocked(self): with pytest.raises(ValueError): _guard_mutation_sql("REVOKE SELECT ON t FROM ROLE r") def test_multi_statement_rejected(self): with pytest.raises(ValueError, match="Multi-statement"): _guard_mutation_sql( "INSERT INTO t VALUES(1); DROP TABLE t" ) def test_empty_raises(self): with pytest.raises(ValueError, match="Empty"): _guard_mutation_sql("") def test_comment_bypass_blocked(self): with pytest.raises(ValueError): _guard_mutation_sql( "/* INSERT */ DROP TABLE t" ) # --------------------------------------------------------------------------- # _to_markdown # --------------------------------------------------------------------------- class TestToMarkdown(unittest.TestCase): def test_empty_rows(self): assert _to_markdown(["col"], []) == "_No rows returned._" def test_single_row(self): out = _to_markdown(["id", "name"], [(1, "Alice")]) assert "| id" in out assert "Alice" in out assert "_1 row_" in out def test_plural_rows(self): rows = [(i,) for i in range(3)] out = _to_markdown(["n"], rows) assert "_3 rows_" in out def test_separator_present(self): out = _to_markdown(["a", "b"], [(1, 2)]) lines = out.splitlines() # second line should be the separator assert all(c in "-| " for c in lines[1]) def test_output_truncated(self): # generate enough rows to exceed MAX_OUTPUT_CHARS big_row = ("x" * 1000,) rows = [big_row] * 200 out = _to_markdown(["col"], rows) assert "truncated" in out assert len(out) <= MAX_OUTPUT_CHARS + 100 # some slack for suffix def test_column_width_matches_longest_value(self): out = _to_markdown( ["short"], [("a very long value indeed",)] ) assert "a very long value indeed" in out # --------------------------------------------------------------------------- # _get_config # --------------------------------------------------------------------------- class TestGetConfig(unittest.TestCase): def _call(self, key, env_val=None, dotenv_val=None): env = {key: env_val} if env_val is not None else {} dotenv = {key: dotenv_val} if dotenv_val is not None else {} fake_path = Path("/fake/.env") if dotenv else None with patch.dict(os.environ, env, clear=True): with patch.object( server, '_find_dotenv', return_value=fake_path ): with patch.object( server, 'dotenv_values', return_value=dotenv ): return _get_config(key) def test_returns_env_var(self): assert self._call("K", env_val="val") == "val" def test_strips_whitespace(self): assert self._call("K", env_val=" val ") == "val" def test_blank_env_uses_dotenv(self): assert self._call( "K", env_val="", dotenv_val="dotval" ) == "dotval" def test_whitespace_env_uses_dotenv(self): assert self._call( "K", env_val=" ", dotenv_val="dotval" ) == "dotval" def test_env_wins_over_dotenv(self): assert self._call( "K", env_val="envval", dotenv_val="dotval" ) == "envval" def test_both_absent_returns_none(self): assert self._call("K") is None def test_both_blank_returns_none(self): assert self._call( "K", env_val="", dotenv_val="" ) is None # --------------------------------------------------------------------------- # _build_connect_kwargs # --------------------------------------------------------------------------- REQUIRED_ENV = { "SNOWFLAKE_ACCOUNT": "myorg-myacct", "SNOWFLAKE_USER": "user@example.com", "SNOWFLAKE_WAREHOUSE": "WH", "SNOWFLAKE_ROLE": "ANALYST", "SNOWFLAKE_DATABASE": "MYDB", "SNOWFLAKE_SCHEMA": "PUBLIC", } class TestBuildConnectKwargs(unittest.TestCase): def _run(self, overrides=None, dotenv_fallback=None): """Run _build_connect_kwargs with a controlled env. dotenv_fallback patches server.dotenv_values so tests are isolated from any real .env file on disk. """ env = {**REQUIRED_ENV, **(overrides or {})} _dotenv = dotenv_fallback if dotenv_fallback is not None else {} fake_path = Path("/fake/.env") if _dotenv else None with patch.dict(os.environ, env, clear=True): with patch.object( server, '_find_dotenv', return_value=fake_path ): with patch.object( server, 'dotenv_values', return_value=_dotenv ): return _build_connect_kwargs() def test_required_fields_present(self): kw = self._run() for key in ( "account", "user", "warehouse", "role", "database", "schema", ): assert key in kw def test_authenticator_is_externalbrowser(self): kw = self._run() assert kw["authenticator"] == "externalbrowser" def test_password_never_in_kwargs(self): kw = self._run({"SNOWFLAKE_PASSWORD": "secret"}) assert "password" not in kw def test_private_key_never_in_kwargs(self): kw = self._run( {"SNOWFLAKE_PRIVATE_KEY_PATH": "/some/key.p8"} ) assert "private_key" not in kw assert "private_key_path" not in kw def test_default_browser_timeout(self): kw = self._run() assert kw["external_browser_timeout"] == 120 def test_custom_browser_timeout(self): kw = self._run({"SNOWFLAKE_BROWSER_TIMEOUT": "60"}) assert kw["external_browser_timeout"] == 60 def test_blank_timeout_falls_back_to_default(self): kw = self._run({"SNOWFLAKE_BROWSER_TIMEOUT": ""}) assert kw["external_browser_timeout"] == 120 def test_missing_account_raises(self): env = {k: v for k, v in REQUIRED_ENV.items() if k != "SNOWFLAKE_ACCOUNT"} with patch.dict(os.environ, env, clear=True): with patch.object(server, '_find_dotenv', return_value=None): with pytest.raises( RuntimeError, match="SNOWFLAKE_ACCOUNT" ): _build_connect_kwargs() def test_missing_user_raises(self): env = {k: v for k, v in REQUIRED_ENV.items() if k != "SNOWFLAKE_USER"} with patch.dict(os.environ, env, clear=True): with patch.object(server, '_find_dotenv', return_value=None): with pytest.raises( RuntimeError, match="SNOWFLAKE_USER" ): _build_connect_kwargs() def test_missing_warehouse_raises(self): env = {k: v for k, v in REQUIRED_ENV.items() if k != "SNOWFLAKE_WAREHOUSE"} with patch.dict(os.environ, env, clear=True): with patch.object(server, '_find_dotenv', return_value=None): with pytest.raises( RuntimeError, match="SNOWFLAKE_WAREHOUSE" ): _build_connect_kwargs() def test_missing_role_raises(self): env = {k: v for k, v in REQUIRED_ENV.items() if k != "SNOWFLAKE_ROLE"} with patch.dict(os.environ, env, clear=True): with patch.object(server, '_find_dotenv', return_value=None): with pytest.raises( RuntimeError, match="SNOWFLAKE_ROLE" ): _build_connect_kwargs() def test_missing_database_raises(self): env = {k: v for k, v in REQUIRED_ENV.items() if k != "SNOWFLAKE_DATABASE"} with patch.dict(os.environ, env, clear=True): with patch.object(server, '_find_dotenv', return_value=None): with pytest.raises( RuntimeError, match="SNOWFLAKE_DATABASE" ): _build_connect_kwargs() def test_missing_schema_raises(self): env = {k: v for k, v in REQUIRED_ENV.items() if k != "SNOWFLAKE_SCHEMA"} with patch.dict(os.environ, env, clear=True): with patch.object(server, '_find_dotenv', return_value=None): with pytest.raises( RuntimeError, match="SNOWFLAKE_SCHEMA" ): _build_connect_kwargs() def test_multiple_missing_all_named(self): with patch.dict(os.environ, {}, clear=True): with patch.object(server, '_find_dotenv', return_value=None): with pytest.raises(RuntimeError) as exc: _build_connect_kwargs() msg = str(exc.value) for var in REQUIRED_ENV: assert var in msg def test_blank_string_treated_as_missing(self): # blank env var with no dotenv fallback → error with pytest.raises( RuntimeError, match="SNOWFLAKE_ACCOUNT" ): self._run( overrides={"SNOWFLAKE_ACCOUNT": " "}, dotenv_fallback={}, ) # --- fallback behaviour --- def test_blank_env_falls_back_to_dotenv(self): # VS Code sends blank; .env has the real value kw = self._run( overrides={"SNOWFLAKE_ACCOUNT": ""}, dotenv_fallback={"SNOWFLAKE_ACCOUNT": "dotenv-acct"}, ) assert kw["account"] == "dotenv-acct" def test_whitespace_env_falls_back_to_dotenv(self): kw = self._run( overrides={"SNOWFLAKE_ACCOUNT": " "}, dotenv_fallback={"SNOWFLAKE_ACCOUNT": "dotenv-acct"}, ) assert kw["account"] == "dotenv-acct" def test_env_var_takes_priority_over_dotenv(self): # non-blank env var wins over .env value kw = self._run( overrides={"SNOWFLAKE_ACCOUNT": "env-acct"}, dotenv_fallback={"SNOWFLAKE_ACCOUNT": "dotenv-acct"}, ) assert kw["account"] == "env-acct" def test_all_vars_can_come_from_dotenv(self): # all required vars absent from os.environ, all in .env dotenv = { "SNOWFLAKE_ACCOUNT": "dot-account", "SNOWFLAKE_USER": "dot-user", "SNOWFLAKE_WAREHOUSE": "dot-wh", "SNOWFLAKE_ROLE": "dot-role", "SNOWFLAKE_DATABASE": "dot-db", "SNOWFLAKE_SCHEMA": "dot-schema", } with patch.dict(os.environ, {}, clear=True): fake_path = Path("/fake/.env") with patch.object( server, '_find_dotenv', return_value=fake_path ): with patch.object( server, 'dotenv_values', return_value=dotenv ): kw = _build_connect_kwargs() assert kw["account"] == "dot-account" assert kw["schema"] == "dot-schema" def test_blank_timeout_falls_back_to_dotenv(self): kw = self._run( overrides={"SNOWFLAKE_BROWSER_TIMEOUT": ""}, dotenv_fallback={ "SNOWFLAKE_BROWSER_TIMEOUT": "60" }, ) assert kw["external_browser_timeout"] == 60 # --------------------------------------------------------------------------- # run_query (tool) # --------------------------------------------------------------------------- class TestRunQuery(unittest.TestCase): def test_select_returns_markdown(self): ctx = _make_ctx(["id", "name"], [(1, "Alice")]) result = run_query("SELECT id, name FROM t", ctx) assert "Alice" in result assert "id" in result def test_limit_injected_on_bare_select(self): ctx = _make_ctx(["n"], [(1,)]) run_query("SELECT 1", ctx) executed_sql = ( ctx.request_context.lifespan_context .conn.cursor.return_value.execute.call_args[0][0] ) assert "LIMIT" in executed_sql.upper() def test_existing_limit_not_doubled(self): ctx = _make_ctx(["n"], [(1,)]) run_query("SELECT 1 LIMIT 5", ctx) executed_sql = ( ctx.request_context.lifespan_context .conn.cursor.return_value.execute.call_args[0][0] ) assert executed_sql.upper().count("LIMIT") == 1 def test_limit_clamped_to_max_rows(self): ctx = _make_ctx(["n"], [(1,)]) run_query("SELECT 1", ctx, limit=9999) executed_sql = ( ctx.request_context.lifespan_context .conn.cursor.return_value.execute.call_args[0][0] ) assert f"LIMIT {MAX_ROWS}" in executed_sql.upper() def test_limit_clamped_minimum_one(self): ctx = _make_ctx(["n"], [(1,)]) run_query("SELECT 1", ctx, limit=0) executed_sql = ( ctx.request_context.lifespan_context .conn.cursor.return_value.execute.call_args[0][0] ) assert "LIMIT 1" in executed_sql.upper() def test_blocked_statement_returns_error_prefix(self): ctx = _make_ctx([], []) result = run_query("DROP TABLE t", ctx) assert result.startswith("⛔") def test_insert_blocked(self): ctx = _make_ctx([], []) result = run_query("INSERT INTO t VALUES(1)", ctx) assert result.startswith("⛔") def test_programming_error_returns_error_prefix(self): ctx = _make_ctx(["n"], []) err = server.snowflake.connector.errors.ProgrammingError( msg="syntax error" ) ctx.request_context.lifespan_context.conn\ .cursor.return_value.execute.side_effect = err result = run_query("SELECT 1", ctx) assert result.startswith("❌") assert "syntax error" in result def test_show_not_limited(self): ctx = _make_ctx(["name"], [("DB1",)]) run_query("SHOW DATABASES", ctx) executed_sql = ( ctx.request_context.lifespan_context .conn.cursor.return_value.execute.call_args[0][0] ) assert "LIMIT" not in executed_sql.upper() def test_cap_banner_shown_at_max_rows(self): rows = [(i,) for i in range(MAX_ROWS)] ctx = _make_ctx(["n"], rows) result = run_query("SELECT n FROM t", ctx) assert "capped" in result.lower() def test_no_cap_banner_below_max_rows(self): ctx = _make_ctx(["n"], [(1,), (2,)]) result = run_query("SELECT n FROM t", ctx) assert "capped" not in result.lower() def test_database_param_passed_to_execute(self): ctx = _make_ctx(["n"], [(1,)]) ret = (["n"], [(1,)]) with patch.object(server, '_execute', return_value=ret) as m: run_query("SELECT 1", ctx, database="MYDB") assert m.call_args[0][2] == "MYDB" assert m.call_args[0][3] is None def test_schema_param_passed_to_execute(self): ctx = _make_ctx(["n"], [(1,)]) ret = (["n"], [(1,)]) with patch.object(server, '_execute', return_value=ret) as m: run_query("SELECT 1", ctx, schema="RAW") assert m.call_args[0][3] == "RAW" def test_invalid_database_rejected(self): ctx = _make_ctx([], []) result = run_query("SELECT 1", ctx, database="bad db!") assert result.startswith("⛔") def test_invalid_schema_rejected(self): ctx = _make_ctx([], []) result = run_query("SELECT 1", ctx, schema="bad schema!") assert result.startswith("⛔") # --------------------------------------------------------------------------- # list_databases # --------------------------------------------------------------------------- class TestListDatabases(unittest.TestCase): def test_name_and_owner_columns(self): ctx = _make_ctx( ["created_on", "name", "owner"], [("2024-01-01", "MYDB", "SYSADMIN")], ) result = list_databases(ctx) assert "MYDB" in result assert "SYSADMIN" in result def test_name_only_when_no_owner_col(self): ctx = _make_ctx( ["created_on", "name"], [("2024-01-01", "MYDB")], ) result = list_databases(ctx) assert "MYDB" in result def test_programming_error(self): ctx = _make_ctx(["name"], []) err = server.snowflake.connector.errors.ProgrammingError( msg="access denied" ) ctx.request_context.lifespan_context.conn\ .cursor.return_value.execute.side_effect = err result = list_databases(ctx) assert result.startswith("❌") # --------------------------------------------------------------------------- # list_schemas # --------------------------------------------------------------------------- class TestListSchemas(unittest.TestCase): def test_returns_schema_names(self): ctx = _make_ctx( ["created_on", "name"], [("2024-01-01", "PUBLIC"), ("2024-01-01", "RAW")], ) result = list_schemas("MYDB", ctx) assert "PUBLIC" in result assert "RAW" in result def test_invalid_database_name(self): ctx = _make_ctx([], []) result = list_schemas("bad name!", ctx) assert result.startswith("⛔") def test_sql_injected_db_rejected(self): ctx = _make_ctx([], []) result = list_schemas("x; DROP TABLE t--", ctx) assert result.startswith("⛔") def test_programming_error(self): ctx = _make_ctx(["name"], []) err = server.snowflake.connector.errors.ProgrammingError( msg="db not found" ) ctx.request_context.lifespan_context.conn\ .cursor.return_value.execute.side_effect = err result = list_schemas("MYDB", ctx) assert result.startswith("❌") # --------------------------------------------------------------------------- # list_tables # --------------------------------------------------------------------------- class TestListTables(unittest.TestCase): def test_returns_table_and_rows_columns(self): ctx = _make_ctx( ["created_on", "name", "rows"], [("2024-01-01", "ORDERS", 1000)], ) result = list_tables("MYDB", "PUBLIC", ctx) assert "ORDERS" in result assert "1000" in result def test_returns_table_only_when_no_rows_col(self): ctx = _make_ctx( ["created_on", "name"], [("2024-01-01", "ORDERS")], ) result = list_tables("MYDB", "PUBLIC", ctx) assert "ORDERS" in result def test_invalid_schema_name(self): ctx = _make_ctx([], []) result = list_tables("MYDB", "bad schema!", ctx) assert result.startswith("⛔") def test_invalid_database_name(self): ctx = _make_ctx([], []) result = list_tables("bad db!", "PUBLIC", ctx) assert result.startswith("⛔") def test_programming_error(self): ctx = _make_ctx(["name"], []) err = server.snowflake.connector.errors.ProgrammingError( msg="schema not found" ) ctx.request_context.lifespan_context.conn\ .cursor.return_value.execute.side_effect = err result = list_tables("MYDB", "PUBLIC", ctx) assert result.startswith("❌") # --------------------------------------------------------------------------- # describe_table # --------------------------------------------------------------------------- class TestDescribeTable(unittest.TestCase): def test_returns_column_info(self): ctx = _make_ctx( ["name", "type", "nullable"], [("ID", "NUMBER", "N"), ("NAME", "TEXT", "Y")], ) result = describe_table("MYDB.PUBLIC.ORDERS", ctx) assert "ID" in result assert "NUMBER" in result def test_invalid_table_name_rejected(self): ctx = _make_ctx([], []) result = describe_table("x; DROP TABLE t", ctx) assert result.startswith("⛔") def test_unqualified_name_ok(self): ctx = _make_ctx( ["name", "type"], [("ID", "NUMBER")] ) result = describe_table("ORDERS", ctx) assert "ID" in result def test_programming_error(self): ctx = _make_ctx(["name"], []) err = server.snowflake.connector.errors.ProgrammingError( msg="table not found" ) ctx.request_context.lifespan_context.conn\ .cursor.return_value.execute.side_effect = err result = describe_table("MYDB.PUBLIC.ORDERS", ctx) assert result.startswith("❌") # --------------------------------------------------------------------------- # preview_mutation / confirm_mutation # --------------------------------------------------------------------------- class TestMutationFlow(unittest.TestCase): def setUp(self): _PENDING_MUTATIONS.clear() def tearDown(self): _PENDING_MUTATIONS.clear() def _preview(self, sql): ctx = _make_ctx([], []) return preview_mutation(sql, ctx) def _confirm(self, token): ctx = _make_ctx_dml() return confirm_mutation(token, ctx) # --- preview --- def test_preview_returns_token(self): result = self._preview("INSERT INTO t VALUES(1)") assert "Approval token" in result assert len(_PENDING_MUTATIONS) == 1 def test_preview_shows_sql(self): sql = "UPDATE t SET x=1 WHERE id=2" result = self._preview(sql) assert sql in result def test_preview_blocked_sql_returns_error(self): result = self._preview("DROP TABLE t") assert result.startswith("⛔") assert len(_PENDING_MUTATIONS) == 0 def test_preview_select_blocked(self): result = self._preview("SELECT 1") assert result.startswith("⛔") # --- confirm --- def test_full_flow_succeeds(self): preview_result = self._preview("INSERT INTO t VALUES(1)") # extract token from between backticks token = preview_result.split("`")[ [i for i, p in enumerate(preview_result.split("`")) if len(p) == 32][0] ] # use the stored token directly instead token = next(iter(_PENDING_MUTATIONS)) result = self._confirm(token) assert result.startswith("✅") def test_token_is_single_use(self): self._preview("DELETE FROM t WHERE id=1") token = next(iter(_PENDING_MUTATIONS)) self._confirm(token) # second use must fail ctx = _make_ctx_dml() result = confirm_mutation(token, ctx) assert result.startswith("⛔") def test_invalid_token_rejected(self): result = self._confirm("not-a-real-token") assert result.startswith("⛔") def test_expired_token_rejected(self): self._preview("INSERT INTO t VALUES(1)") token = next(iter(_PENDING_MUTATIONS)) # back-date the expiry sql, _, db, sc = _PENDING_MUTATIONS[token] _PENDING_MUTATIONS[token] = (sql, time.time() - 1, db, sc) ctx = _make_ctx_dml() result = confirm_mutation(token, ctx) assert result.startswith("⛔") assert "expired" in result.lower() def test_stale_tokens_purged_on_preview(self): # plant two expired tokens (full 4-tuple) _PENDING_MUTATIONS["old1"] = ("X", time.time() - 1, None, None) _PENDING_MUTATIONS["old2"] = ("Y", time.time() - 1, None, None) self._preview("INSERT INTO t VALUES(1)") assert "old1" not in _PENDING_MUTATIONS assert "old2" not in _PENDING_MUTATIONS def test_confirm_programming_error(self): self._preview("INSERT INTO t VALUES(1)") token = next(iter(_PENDING_MUTATIONS)) ctx = _make_ctx_dml() err = server.snowflake.connector.errors.ProgrammingError( msg="constraint violation" ) ctx.request_context.lifespan_context.conn\ .cursor.return_value.execute.side_effect = err result = confirm_mutation(token, ctx) assert result.startswith("❌") assert "constraint violation" in result def test_token_ttl_is_set_correctly(self): self._preview("INSERT INTO t VALUES(1)") token = next(iter(_PENDING_MUTATIONS)) _, expires_at, _db, _sc = _PENDING_MUTATIONS[token] # Should expire roughly MUTATION_TOKEN_TTL seconds from now assert abs(expires_at - time.time() - MUTATION_TOKEN_TTL) < 2 def test_create_table_allowed_in_mutation(self): result = self._preview("CREATE TABLE new_t (id INT)") assert "Approval token" in result def test_create_schema_blocked_in_mutation(self): result = self._preview("CREATE SCHEMA s") assert result.startswith("⛔") def test_context_stored_in_token(self): """database/schema passed to preview_mutation are stored in the pending mutations dict and visible in the output.""" ctx = _make_ctx([], []) preview_mutation( "INSERT INTO t VALUES(1)", ctx, database="MYDB", schema="RAW", ) token = next(iter(_PENDING_MUTATIONS)) _, _exp, db, sc = _PENDING_MUTATIONS[token] assert db == "MYDB" assert sc == "RAW" def test_context_note_shown_in_preview(self): ctx = _make_ctx([], []) result = preview_mutation( "INSERT INTO t VALUES(1)", ctx, database="MYDB", schema="RAW", ) assert "MYDB" in result assert "RAW" in result def test_invalid_database_in_preview_rejected(self): ctx = _make_ctx([], []) result = preview_mutation( "INSERT INTO t VALUES(1)", ctx, database="bad db!", ) assert result.startswith("⛔") assert len(_PENDING_MUTATIONS) == 0 def test_invalid_schema_in_preview_rejected(self): ctx = _make_ctx([], []) result = preview_mutation( "INSERT INTO t VALUES(1)", ctx, schema="bad schema!", ) assert result.startswith("⛔") assert len(_PENDING_MUTATIONS) == 0 # --------------------------------------------------------------------------- # Graceful degradation — tools return ⚙️ message on RuntimeError # --------------------------------------------------------------------------- class TestGracefulDegradation(unittest.TestCase): """Verify every tool that calls _execute catches RuntimeError and returns a user-facing configuration error string instead of propagating the exception.""" _CONFIG_ERR = "missing creds" def _ctx(self): return _make_ctx(["n"], [(1,)]) def _raise_runtime(self, *_a, **_kw): raise RuntimeError(self._CONFIG_ERR) # --- run_query --- def test_run_query_config_error(self): ctx = self._ctx() with patch.object(server, '_execute', self._raise_runtime): result = run_query("SELECT 1", ctx) assert result.startswith("⚙️") assert self._CONFIG_ERR in result # --- list_databases --- def test_list_databases_config_error(self): ctx = self._ctx() with patch.object(server, '_execute', self._raise_runtime): result = list_databases(ctx) assert result.startswith("⚙️") assert self._CONFIG_ERR in result # --- list_schemas --- def test_list_schemas_config_error(self): ctx = self._ctx() with patch.object(server, '_execute', self._raise_runtime): result = list_schemas("MYDB", ctx) assert result.startswith("⚙️") assert self._CONFIG_ERR in result # --- list_tables --- def test_list_tables_config_error(self): ctx = self._ctx() with patch.object(server, '_execute', self._raise_runtime): result = list_tables("MYDB", "PUBLIC", ctx) assert result.startswith("⚙️") assert self._CONFIG_ERR in result # --- describe_table --- def test_describe_table_config_error(self): ctx = self._ctx() with patch.object(server, '_execute', self._raise_runtime): result = describe_table("MYDB.PUBLIC.ORDERS", ctx) assert result.startswith("⚙️") assert self._CONFIG_ERR in result # --- confirm_mutation --- def test_confirm_mutation_config_error(self): _PENDING_MUTATIONS.clear() preview_ctx = _make_ctx([], []) preview_mutation("INSERT INTO t VALUES(1)", preview_ctx) token = next(iter(_PENDING_MUTATIONS)) ctx = self._ctx() with patch.object(server, '_execute', self._raise_runtime): result = confirm_mutation(token, ctx) assert result.startswith("⚙️") assert self._CONFIG_ERR in result _PENDING_MUTATIONS.clear()