using CareFix.Core; using Xunit; namespace CareFix.Tests; /// The guard is the first line of defence on every query the AI writes. public class SqlGuardTests { private readonly SqlGuard _guard = new(); private static readonly HashSet Known = new(StringComparer.OrdinalIgnoreCase) { "IPD_BILL", "IPD_PHARMACY_ISSUE", "PATIENT_MASTER" }; private bool Allowed(string sql) => _guard.ValidateSelect(sql, Known).Ok; [Theory] [InlineData("SELECT BILL_ID, NET_AMT FROM IPD_BILL WHERE BILL_ID = 77431")] [InlineData("SELECT TOP 10 * FROM dbo.IPD_BILL ORDER BY BILL_ID DESC")] [InlineData("SELECT b.BILL_ID, i.QTY FROM IPD_BILL b JOIN IPD_PHARMACY_ISSUE i ON i.IPD_NO = b.IPD_NO WHERE b.BILL_ID = 1")] [InlineData("WITH x AS (SELECT IPD_NO, SUM(AMOUNT) total FROM IPD_PHARMACY_ISSUE GROUP BY IPD_NO) SELECT * FROM x WHERE total > 100")] [InlineData("SELECT COUNT(*) FROM IPD_BILL WITH (NOLOCK)")] public void Allows_plain_reads(string sql) => Assert.True(Allowed(sql), sql); [Theory] [InlineData("UPDATE IPD_BILL SET NET_AMT = 0")] [InlineData("DELETE FROM IPD_BILL WHERE BILL_ID = 1")] [InlineData("DROP TABLE IPD_BILL")] [InlineData("TRUNCATE TABLE IPD_BILL")] [InlineData("EXEC sp_who")] [InlineData("SELECT 1; DROP TABLE IPD_BILL")] // stacked statements [InlineData("SELECT * INTO copy_of_bills FROM IPD_BILL")] // SELECT INTO writes [InlineData("SELECT * FROM OtherDb.dbo.IPD_BILL")] // cross database [InlineData("SELECT * FROM LinkedServer.OtherDb.dbo.IPD_BILL")] // linked server [InlineData("SELECT * FROM OPENROWSET('SQLNCLI', 'x', 'SELECT 1')")] [InlineData("SELECT * FROM OPENQUERY(remote, 'SELECT 1')")] [InlineData("SELECT * FROM sys.objects")] // non-dbo schema [InlineData("SELECT * FROM INFORMATION_SCHEMA.COLUMNS")] [InlineData("SELECT * FROM SECRET_SALARY_TABLE")] // not in this hospital's schema [InlineData("SELECT * FROM dbo.fn_Something(1)")] // table-valued function [InlineData("this is not sql")] [InlineData("")] public void Blocks_everything_else(string sql) => Assert.False(Allowed(sql), sql); [Fact] public void Reports_the_unknown_table_by_name() { var r = _guard.ValidateSelect("SELECT * FROM PAYROLL_MASTER", Known); Assert.False(r.Ok); Assert.Contains("PAYROLL_MASTER", r.Reason); } [Fact] public void Lists_the_tables_a_query_touches() { var r = _guard.ValidateSelect("SELECT b.BILL_ID FROM IPD_BILL b JOIN IPD_PHARMACY_ISSUE i ON i.IPD_NO = b.IPD_NO", Known); Assert.True(r.Ok); Assert.Equal(2, r.Tables.Count); } [Fact] public void Skips_the_table_check_when_no_schema_is_given() { // Used when a rule or playbook query is imported, before any hospital is chosen. Assert.True(_guard.ValidateSelect("SELECT 1 FROM ANY_TABLE WHERE ID = @PkValue", null).Ok); Assert.False(_guard.ValidateSelect("UPDATE ANY_TABLE SET X = 1", null).Ok); } }