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);
}
}