use serde_json::{Map, Value};
use tiberius::ToSql;
use crate::auth::auth_manager::AuthManager;
use crate::core::config_schema::Config;
use crate::data::store::EndpointRecord;
use crate::services::sql_pool;
use crate::services::sql_type::{column_data_to_json, json_to_param};
use crate::validation::validator::resolved_schemas_for;
fn parse_path(path: &str) -> anyhow::Result<(&str, &str, &str)> {
let mut parts = path.trim_start_matches('/').splitn(3, '/');
match (parts.next(), parts.next(), parts.next()) {
(Some(db), Some(schema), Some(name))
if !db.is_empty() && !schema.is_empty() && !name.is_empty() =>
{
Ok((
validate_ident(db)?,
validate_ident(schema)?,
validate_ident(name)?,
))
}
_ => anyhow::bail!(
"endpoint path '{path}' is not in the expected /<db>/<schema>/<name> shape"
),
}
}
fn validate_ident(ident: &str) -> anyhow::Result<&str> {
if !ident.is_empty()
&& ident
.bytes()
.all(|b| b.is_ascii_alphanumeric() || b == b'_')
{
Ok(ident)
} else {
anyhow::bail!("identifier '{ident}' contains characters outside [A-Za-z0-9_]")
}
}
fn quote_ident(ident: &str) -> String {
format!("[{}]", ident.replace(']', "]]"))
}
struct Param {
name: String,
ordinal: u64,
x_sql_type: String,
}
fn ordered_params(input_schema: &Value) -> Vec<Param> {
let mut params: Vec<Param> = input_schema
.get("properties")
.and_then(Value::as_object)
.into_iter()
.flatten()
.map(|(name, schema)| Param {
name: name.clone(),
ordinal: schema
.get("x-sql-ordinal")
.and_then(Value::as_u64)
.unwrap_or(0),
x_sql_type: schema
.get("x-sql-type")
.and_then(Value::as_str)
.unwrap_or("nvarchar(max)")
.to_string(),
})
.collect();
params.sort_by_key(|p| p.ordinal);
params
}
fn build_statement(
db: &str,
schema: &str,
name: &str,
kind: Option<&str>,
params: &[Param],
body: &Map<String, Value>,
database_override: Option<&str>,
) -> anyhow::Result<(String, Vec<Box<dyn ToSql>>)> {
let qualified = match (db, database_override) {
("sandbox", Some(requested_db)) => format!(
"{}.{}.{}",
quote_ident(requested_db),
quote_ident(schema),
quote_ident(name)
),
("sandbox", None) => format!("{}.{}", quote_ident(schema), quote_ident(name)),
(_, _) => format!(
"{}.{}.{}",
quote_ident(db),
quote_ident(schema),
quote_ident(name)
),
};
let mut bound: Vec<Box<dyn ToSql>> = Vec::with_capacity(params.len());
for param in params {
let value = body.get(¶m.name).cloned().unwrap_or(Value::Null);
bound.push(json_to_param(&value, ¶m.x_sql_type)?);
}
let sql = match kind {
Some("VIEW") => format!("SELECT * FROM {qualified}"),
Some(k) if k.ends_with("_FUNCTION") || k.contains("TABLE_VALUED_FUNCTION") => {
let placeholders = (1..=bound.len())
.map(|i| format!("@P{i}"))
.collect::<Vec<_>>()
.join(", ");
format!("SELECT * FROM {qualified}({placeholders})")
}
_ => {
if bound.is_empty() {
format!("EXEC {qualified}")
} else {
let mut assignments = Vec::with_capacity(params.len());
for (i, p) in params.iter().enumerate() {
assignments.push(format!(
"{} = @P{}",
quote_ident(validate_ident(&p.name)?),
i + 1
));
}
format!("EXEC {qualified} {}", assignments.join(", "))
}
}
};
Ok((sql, bound))
}
fn classify_tiberius_error(err: tiberius::error::Error) -> anyhow::Error {
use tiberius::error::Error as TError;
match &err {
TError::Server(token) => {
let (number, state, class, message, procedure, line) = (
token.code(),
token.state(),
token.class(),
token.message(),
token.procedure(),
token.line(),
);
let category = if class >= 17 {
"fatal/resource error (500-equivalent)"
} else if class == 14 {
"permission denied (403-equivalent)"
} else {
"statement/user error (400-equivalent)"
};
anyhow::anyhow!(
"SQL Server error {number} (severity {class}, state {state}) in '{procedure}' line {line}: {message} [{category}]"
)
}
other => anyhow::anyhow!("SQL Server connection/protocol error (500-equivalent): {other}"),
}
}
pub struct ApiClient {
config: Config,
}
impl ApiClient {
pub fn new(config: Config) -> Self {
Self { config }
}
pub async fn execute(
&self,
endpoint: &EndpointRecord,
args: &Value,
auth_manager: &mut AuthManager,
) -> anyhow::Result<Value> {
let (db, schema, name) = parse_path(&endpoint.path)?;
let (input_schema, _output_schema) =
resolved_schemas_for(&self.config.api_version, &endpoint.operation_id).ok_or_else(
|| {
anyhow::anyhow!(
"no resolved schema found for operation '{}' under api_version '{}'",
endpoint.operation_id,
self.config.api_version
)
},
)?;
let params = ordered_params(input_schema);
let empty = Map::new();
let args_map = args.as_object().unwrap_or(&empty);
let body = args_map
.get("body")
.and_then(Value::as_object)
.cloned()
.unwrap_or_default();
let database_override = args_map
.get("database")
.and_then(Value::as_str)
.map(validate_ident)
.transpose()?;
let (sql, bound) = build_statement(
db,
schema,
name,
endpoint.description.as_deref(),
¶ms,
&body,
database_override,
)?;
let auth_method = auth_manager.resolve_tds_auth().await?;
let (host, port) = self.config.host_and_port();
let mut tiberius_config = tiberius::Config::new();
tiberius_config.host(host);
tiberius_config.port(port);
tiberius_config.authentication(auth_method);
if self.config.trust_server_cert {
tiberius_config.trust_cert();
}
let pool_key = format!("{host}:{port}");
let pool =
sql_pool::cached_pool(&pool_key, tiberius_config, self.config.pool_max_size).await?;
let mut conn = pool.get().await.map_err(|err| {
anyhow::anyhow!("failed to obtain a pooled SQL Server connection: {err}")
})?;
let bound_refs: Vec<&dyn ToSql> = bound.iter().map(|param| param.as_ref()).collect();
let stream = conn
.query(&sql, &bound_refs)
.await
.map_err(classify_tiberius_error)?;
let rows = stream
.into_first_result()
.await
.map_err(classify_tiberius_error)?;
let json_rows: Vec<Value> = rows
.iter()
.map(|row| {
let mut obj = Map::new();
for (column, data) in row.cells() {
obj.insert(column.name().to_string(), column_data_to_json(data));
}
Value::Object(obj)
})
.collect();
Ok(Value::Array(json_rows))
}
}
#[cfg(test)]
mod tests {
use super::*;
#[test]
fn parse_path_splits_db_schema_name() {
assert_eq!(
parse_path("/master/INFORMATION_SCHEMA/COLUMNS").unwrap(),
("master", "INFORMATION_SCHEMA", "COLUMNS")
);
}
#[test]
fn parse_path_rejects_a_path_with_too_few_segments() {
assert!(parse_path("/master/COLUMNS").is_err());
}
#[test]
fn quote_ident_doubles_embedded_closing_brackets() {
assert_eq!(quote_ident("weird]name"), "[weird]]name]");
}
#[test]
fn validate_ident_accepts_alphanumeric_and_underscore() {
assert!(validate_ident("sp_who2").is_ok());
assert!(validate_ident("INFORMATION_SCHEMA").is_ok());
}
#[test]
fn validate_ident_rejects_anything_else() {
assert!(validate_ident("").is_err());
assert!(validate_ident("robert'; drop table endpoints;--").is_err());
assert!(validate_ident("weird]name").is_err());
assert!(validate_ident("has space").is_err());
}
#[test]
fn parse_path_rejects_a_path_with_an_unsafe_segment() {
assert!(parse_path("/master/sys/sp_who; DROP TABLE endpoints--").is_err());
}
#[test]
fn build_statement_selects_from_a_view_with_no_parameters() {
let (sql, bound) = build_statement(
"master",
"INFORMATION_SCHEMA",
"COLUMNS",
Some("VIEW"),
&[],
&Map::new(),
None,
)
.unwrap();
assert_eq!(sql, "SELECT * FROM [master].[INFORMATION_SCHEMA].[COLUMNS]");
assert!(bound.is_empty());
}
#[test]
fn build_statement_execs_a_proc_with_named_parameters() {
let params = vec![
Param {
name: "objname".to_string(),
ordinal: 1,
x_sql_type: "nvarchar(1035)".to_string(),
},
Param {
name: "newname".to_string(),
ordinal: 2,
x_sql_type: "sysname".to_string(),
},
];
let mut body = Map::new();
body.insert("objname".to_string(), Value::String("t1".to_string()));
body.insert("newname".to_string(), Value::String("t2".to_string()));
let (sql, bound) = build_statement(
"master",
"sys",
"sp_rename",
Some("SQL_STORED_PROCEDURE"),
¶ms,
&body,
None,
)
.unwrap();
assert_eq!(
sql,
"EXEC [master].[sys].[sp_rename] [objname] = @P1, [newname] = @P2"
);
assert_eq!(bound.len(), 2);
}
#[test]
fn build_statement_selects_from_a_function_with_positional_placeholders() {
let params = vec![Param {
name: "session_id".to_string(),
ordinal: 1,
x_sql_type: "smallint".to_string(),
}];
let mut body = Map::new();
body.insert("session_id".to_string(), Value::from(52));
let (sql, bound) = build_statement(
"master",
"sys",
"dm_exec_sql_text",
Some("SQL_INLINE_TABLE_VALUED_FUNCTION"),
¶ms,
&body,
None,
)
.unwrap();
assert_eq!(sql, "SELECT * FROM [master].[sys].[dm_exec_sql_text](@P1)");
assert_eq!(bound.len(), 1);
}
#[test]
fn build_statement_omits_the_database_qualifier_for_sandbox_by_default() {
let (sql, bound) = build_statement(
"sandbox",
"dbo",
"widgets",
Some("VIEW"),
&[],
&Map::new(),
None,
)
.unwrap();
assert_eq!(sql, "SELECT * FROM [dbo].[widgets]");
assert!(bound.is_empty());
}
#[test]
fn build_statement_replaces_sandbox_with_a_requested_database_override() {
let (sql, bound) = build_statement(
"sandbox",
"dbo",
"widgets",
Some("VIEW"),
&[],
&Map::new(),
Some("reporting"),
)
.unwrap();
assert_eq!(sql, "SELECT * FROM [reporting].[dbo].[widgets]");
assert!(bound.is_empty());
}
#[test]
fn build_statement_ignores_the_database_override_for_master_and_msdb() {
let (sql, _bound) = build_statement(
"master",
"sys",
"sp_who",
Some("SQL_STORED_PROCEDURE"),
&[],
&Map::new(),
Some("reporting"),
)
.unwrap();
assert_eq!(sql, "EXEC [master].[sys].[sp_who]");
}
#[test]
fn build_statement_execs_a_parameterless_proc() {
let (sql, bound) = build_statement(
"master",
"sys",
"sp_who",
Some("SQL_STORED_PROCEDURE"),
&[],
&Map::new(),
None,
)
.unwrap();
assert_eq!(sql, "EXEC [master].[sys].[sp_who]");
assert!(bound.is_empty());
}
#[test]
fn ordered_params_follow_the_schema_ordinals() {
let schema = serde_json::json!({
"properties": {
"second": { "x-sql-ordinal": 2, "x-sql-type": "int" },
"first": { "x-sql-ordinal": 1, "x-sql-type": "nvarchar(20)" }
}
});
let params = ordered_params(&schema);
assert_eq!(params.len(), 2);
assert_eq!(params[0].name, "first");
assert_eq!(params[0].x_sql_type, "nvarchar(20)");
assert_eq!(params[1].name, "second");
}
#[test]
fn build_statement_rejects_an_unsafe_parameter_name() {
let params = vec![Param {
name: "unsafe name".to_string(),
ordinal: 1,
x_sql_type: "int".to_string(),
}];
assert!(
build_statement(
"master",
"sys",
"sp_example",
Some("SQL_STORED_PROCEDURE"),
¶ms,
&Map::new(),
None,
)
.is_err()
);
}
#[test]
fn protocol_errors_are_classified_as_server_failures() {
let error =
classify_tiberius_error(tiberius::error::Error::Protocol("invalid packet".into()));
assert_eq!(
error.to_string(),
"SQL Server connection/protocol error (500-equivalent): Protocol error: invalid packet"
);
}
}