From 7f85659d11233f4bce6d7a96c2f7d767dc97722c Mon Sep 17 00:00:00 2001 From: celsowm <369336+celsowm@users.noreply.github.com> Date: Tue, 21 Apr 2026 15:59:09 +0000 Subject: [PATCH] Refactor iridium_server modules for Single Responsibility Principle - Extracted `execute_sql` and batch logic from `TdsSession` into `execution.rs`. - Extracted handshake and login logic into `handshake.rs`. - Refactored `tds::rpc::parser.rs` into a module directory with `types.rs`, `utils.rs`, and `parser.rs`. - Validated via `cargo test` and `cargo check`. Core parsers in `iridium_core` remain un-split due to tight cross-module trait coupling. --- .../iridium_server/src/session/execution.rs | 235 +++ .../iridium_server/src/session/handshake.rs | 1119 +++++++++++++ crates/iridium_server/src/session/mod.rs | 1407 +---------------- .../iridium_server/src/tds/rpc/parser/mod.rs | 110 ++ .../src/tds/rpc/{ => parser}/parser.rs | 399 +---- .../src/tds/rpc/parser/types.rs | 118 ++ .../src/tds/rpc/parser/utils.rs | 130 ++ 7 files changed, 1755 insertions(+), 1763 deletions(-) create mode 100644 crates/iridium_server/src/session/execution.rs create mode 100644 crates/iridium_server/src/session/handshake.rs create mode 100644 crates/iridium_server/src/tds/rpc/parser/mod.rs rename crates/iridium_server/src/tds/rpc/{ => parser}/parser.rs (64%) create mode 100644 crates/iridium_server/src/tds/rpc/parser/types.rs create mode 100644 crates/iridium_server/src/tds/rpc/parser/utils.rs diff --git a/crates/iridium_server/src/session/execution.rs b/crates/iridium_server/src/session/execution.rs new file mode 100644 index 0000000..35a5aa0 --- /dev/null +++ b/crates/iridium_server/src/session/execution.rs @@ -0,0 +1,235 @@ +use tokio::io::AsyncWriteExt; +use iridium_core::types::{DataType, Value}; +use iridium_core::SessionId; +use super::TdsSession; +use crate::tds::batch::{build_error_response, parse_sql_batch}; +use crate::tds::packet::{self, PacketBuilder, TABULAR_RESULT}; +use crate::tds::tokens; +use crate::session::response::{build_single_int_result, build_use_database_response}; +use crate::session::compat::{extract_leading_use_database, is_ssms_contained_auth_probe, parse_simple_use_database}; + +impl TdsSession { + pub(crate) async fn handle_sql_batch( + &mut self, + data: &[u8], + writer: &mut W, + ) -> Result { + let sql = match parse_sql_batch(data) { + Ok(s) => s, + Err(e) => { + let err = iridium_core::error::DbError::Parse(e.to_string()); + let err_resp = build_error_response(&err); + let _ = packet::write_packet(writer, TABULAR_RESULT, &err_resp.data).await; + return Ok(true); + } + }; + + if !sql.trim().is_empty() { + log::info!( + "[conn={}] SQL batch received:\n{}", + self.connection_id, + crate::session::format_sql_for_log(sql.trim()) + ); + } + self.execute_sql(sql.trim(), writer).await + } + + pub(crate) async fn execute_sql( + &mut self, + sql: &str, + writer: &mut W, + ) -> Result { + if sql.is_empty() { + let mut b = PacketBuilder::new(); + tokens::write_done(&mut b, tokens::DONE_FINAL, 1, 0); + let _ = packet::write_packet(writer, TABULAR_RESULT, b.as_bytes()).await; + return Ok(true); + } + + let session_id = self.session_id.ok_or_else(|| { + iridium_core::error::DbError::Execution("session not initialized".to_string()) + })?; + + if is_ssms_contained_auth_probe(sql) { + if let Some(db_name) = extract_leading_use_database(sql) { + self.apply_use_database(session_id, &db_name, writer) + .await?; + } + let data = build_single_int_result("", 0); + let _ = packet::write_packet(writer, TABULAR_RESULT, &data).await; + return Ok(true); + } + + if let Some(db_name) = parse_simple_use_database(sql) { + self.apply_use_database(session_id, &db_name, writer) + .await?; + return Ok(true); + } + + crate::session::log_sql_execution(self.connection_id, sql); + let force_sysdac_probe_int = crate::session::compat::is_sysdac_instances_probe(sql); + match self.db.execute_session_batch_sql_multi(session_id, sql) { + Ok(results) => { + let count = results.len(); + let mut b = PacketBuilder::with_capacity(4096); + let textsize = self + .db + .session_options(session_id) + .map(|opts| opts.textsize.max(0) as usize) + .unwrap_or(4096); + + for (i, result) in results.into_iter().enumerate() { + let is_last = i == count - 1; + + match result { + Some(mut query_result) => { + let is_proc = query_result.is_procedure; + let return_status = query_result.return_status; + + if !query_result.columns.is_empty() { + if force_sysdac_probe_int + && query_result.columns.len() == 1 + && query_result.rows.len() == 1 + { + query_result.column_types[0] = DataType::Int; + if let Some(row) = query_result.rows.get_mut(0) { + if let Some(value) = row.get_mut(0) { + let int_val = match &*value { + Value::Null => 0, + other => other.to_integer_i64().unwrap_or(0) as i32, + }; + *value = Value::Int(int_val); + } + } + } + + let mut types = Vec::new(); + log::debug!( + "[conn={}] Result set: columns={}, types={}", + self.connection_id, + query_result.columns.len(), + query_result.column_types.len() + ); + for ct in &query_result.column_types { + types.push(crate::tds::type_mapping::runtime_type_to_tds(ct)); + } + for (idx, col_name) in query_result.columns.iter().enumerate() { + if let (Some(runtime_ty), Some(tds_ty)) = + (query_result.column_types.get(idx), types.get(idx)) + { + log::debug!( + "[conn={}] COLMETADATA[{}]: name='{}' runtime={:?} tds=0x{:02X} len={:02X?}", + self.connection_id, + idx, + col_name, + runtime_ty, + tds_ty.tds_type, + tds_ty.length_prefix + ); + } + } + tokens::write_colmetadata(&mut b, &query_result.columns, &types); + for row in &query_result.rows { + tokens::write_row(&mut b, row, &types, textsize); + } + + if is_proc { + tokens::write_done_in_proc( + &mut b, + tokens::DONE_MORE | tokens::DONE_COUNT, + 1, + query_result.rows.len() as u64, + ); + } else { + let done_status = if is_last && return_status.is_none() { + tokens::DONE_FINAL + } else { + tokens::DONE_MORE + }; + tokens::write_done( + &mut b, + done_status | tokens::DONE_COUNT, + 1, + query_result.rows.len() as u64, + ); + } + } else if !is_proc { + let done_status = if is_last && return_status.is_none() { + tokens::DONE_FINAL + } else { + tokens::DONE_MORE + }; + tokens::write_done( + &mut b, + done_status | tokens::DONE_COUNT, + 1, + query_result.rows.len() as u64, + ); + } + + if let Some(code) = return_status { + tokens::write_returnstatus(&mut b, code); + let done_status = if is_last { + tokens::DONE_FINAL + } else { + tokens::DONE_MORE + }; + tokens::write_doneproc(&mut b, done_status, 1, 0); + } + } + None => { + let done_status = if is_last { + tokens::DONE_FINAL + } else { + tokens::DONE_MORE + }; + tokens::write_done(&mut b, done_status, 1, 0); + } + } + } + + if count == 0 { + tokens::write_done(&mut b, tokens::DONE_FINAL, 1, 0); + } + + let _ = packet::write_packet(writer, TABULAR_RESULT, b.as_bytes()).await; + } + Err(e) => { + log::warn!( + "[conn={}] SQL execution failed for batch:\n{}\nerror: {}", + self.connection_id, + crate::session::format_sql_for_log(sql), + e + ); + let err_resp = build_error_response(&e); + let _ = packet::write_packet(writer, TABULAR_RESULT, &err_resp.data).await; + } + } + + Ok(true) + } + + pub(crate) async fn apply_use_database( + &mut self, + session_id: SessionId, + db_name: &str, + writer: &mut W, + ) -> Result<(), iridium_core::error::DbError> { + let old_db = self.database.clone(); + self.database = db_name.to_string(); + if let Err(e) = self + .db + .set_session_database(session_id, self.database.clone()) + { + log::error!( + "[conn={}] Failed to update session database context: {}", + self.connection_id, + e + ); + } + + let data = build_use_database_response(&self.database, &old_db); + let _ = packet::write_packet(writer, TABULAR_RESULT, &data).await; + Ok(()) + } +} diff --git a/crates/iridium_server/src/session/handshake.rs b/crates/iridium_server/src/session/handshake.rs new file mode 100644 index 0000000..150dd80 --- /dev/null +++ b/crates/iridium_server/src/session/handshake.rs @@ -0,0 +1,1119 @@ +use std::sync::Arc; +use std::sync::atomic::{AtomicBool, Ordering}; +use tokio::io::AsyncWriteExt; +use tokio::net::TcpStream; +use iridium_core::error::DbError; +use super::{TdsSession, AsyncReadWrite}; +use crate::tds::batch::build_error_response; +use crate::tds::bulk::parse_bulk_load_data; +use crate::tds::login::parse_login7; +use crate::tds::packet::{ + self, PacketBuilder, ATTENTION, BULK_LOAD, RPC, SQL_BATCH, TABULAR_RESULT, TDS7_LOGIN, + TDS7_PRELOGIN, +}; +use crate::tds::prelogin::{ + build_prelogin_response, parse_prelogin, ENCRYPT_NOT_SUP, ENCRYPT_OFF, ENCRYPT_ON, + ENCRYPT_REQUIRED, +}; +use crate::tds::rpc::{ + parse_param_decl, parse_rpc, + CatalogProc, RpcRequest, CursorOp +}; +use crate::tds::rpc::{ + build_param_preamble, build_param_preamble_with_decls +}; +use crate::tds::tokens; +use crate::tds_tls_io::TdsTlsIo; +use crate::tls; +use crate::pool::CheckoutError; +use super::PreparedStatement; + +impl TdsSession { + pub async fn handle(self, mut stream: TcpStream) -> Result<(), String> { + let mut needs_tls_upgrade = false; + let mut login_packet = None; + + log::info!( + "[conn={}] Starting handshake for incoming connection", + self.connection_id + ); + loop { + let result: Result<_, std::io::Error> = packet::read_message(&mut stream).await; + let (header, data) = result.map_err(|e| format!("Handshake read error: {}", e))?; + crate::session::log_packet(self.connection_id, "handshake", &header, &data); + + if header.packet_type == TDS7_PRELOGIN { + log::debug!( + "[conn={}] PRELOGIN data hex: {:02X?}", + self.connection_id, + data + ); + let prelogin = parse_prelogin(&data).map_err(|e| e.to_string())?; + log::debug!( + "[conn={}] PRELOGIN: version={:?}, encryption={}", + self.connection_id, + prelogin.version, + prelogin.encryption + ); + + let server_encrypt = if self.config.tls_enabled { + ENCRYPT_ON + } else if prelogin.encryption == ENCRYPT_NOT_SUP + || prelogin.encryption == ENCRYPT_ON + || prelogin.encryption == ENCRYPT_REQUIRED + { + ENCRYPT_NOT_SUP + } else { + ENCRYPT_OFF + }; + + needs_tls_upgrade = self.config.tls_enabled + && (prelogin.encryption == ENCRYPT_ON + || prelogin.encryption == ENCRYPT_REQUIRED); + + let response = build_prelogin_response(server_encrypt); + packet::write_packet(&mut stream, TDS7_PRELOGIN, &response) + .await + .map_err(|e| format!("Failed to write PRELOGIN response: {}", e))?; + stream + .flush() + .await + .map_err(|e| format!("Failed to flush: {}", e))?; + log::debug!( + "[conn={}] Sent PRELOGIN response (encryption={})", + self.connection_id, + server_encrypt + ); + + if needs_tls_upgrade { + break; + } + } else if header.packet_type == TDS7_LOGIN { + crate::session::log_packet(self.connection_id, "login", &header, &data); + login_packet = Some((header, data)); + break; + } else { + return Err(format!( + "Expected PRELOGIN or LOGIN7, got 0x{:02X}", + header.packet_type + )); + } + } + + self.handle_login(stream, login_packet, needs_tls_upgrade) + .await + } + + pub(crate) async fn handle_login( + mut self, + stream: TcpStream, + login_packet: Option<(packet::PacketHeader, Vec)>, + needs_tls_upgrade: bool, + ) -> Result<(), String> { + let stream: Box = if needs_tls_upgrade { + log::info!( + "[conn={}] Client requested TLS, upgrading connection via TDS-wrapped handshake", + self.connection_id + ); + + let tls_config = if let (Some(cert_path), Some(key_path)) = + (&self.config.tls_cert_path, &self.config.tls_key_path) + { + tls::load_tls_config(cert_path, key_path) + .map_err(|e| format!("Failed to load TLS config: {}", e))? + } else { + return Err("TLS enabled but no certificate configured".to_string()); + }; + + let acceptor = tokio_rustls::TlsAcceptor::from(std::sync::Arc::new(tls_config)); + + let raw_mode = Arc::new(AtomicBool::new(false)); + let tds_io = TdsTlsIo::new(stream, raw_mode.clone()); + + let tls_stream = acceptor + .accept(tds_io) + .await + .map_err(|e| format!("TLS handshake failed: {}", e))?; + + raw_mode.store(true, Ordering::Release); + + log::info!( + "[conn={}] TLS handshake completed, switched to raw mode", + self.connection_id + ); + Box::new(tls_stream) + } else { + Box::new(stream) + }; + + let (mut reader, mut writer) = tokio::io::split(stream); + + let login_data = if let Some((_, data)) = login_packet { + data + } else { + let result: Result<_, std::io::Error> = packet::read_message(&mut reader).await; + let (header, data) = result.map_err(|e| format!("Failed to read LOGIN7: {}", e))?; + if header.packet_type != TDS7_LOGIN { + return Err(format!( + "Expected LOGIN7 (0x10), got 0x{:02X}", + header.packet_type + )); + } + data + }; + + let login = parse_login7(&login_data).map_err(|e| e.to_string())?; + log::info!( + "[conn={}] LOGIN7: user={}, database={}, app={}, packet_size={}", + self.connection_id, + login.username, + login.database, + login.app_name, + login.packet_size + ); + + if let Some(ref creds) = self.config.auth { + if login.username != creds.user || login.password != creds.password { + log::warn!( + "[conn={}] Login rejected for user={} against configured SQL auth", + self.connection_id, + login.username + ); + let err = DbError::Execution("Login failed for user.".to_string()); + let err_resp = build_error_response(&err); + let _ = packet::write_packet(&mut writer, TABULAR_RESULT, &err_resp.data).await; + return Ok(()); + } + log::info!( + "[conn={}] Login accepted for user={}", + self.connection_id, + login.username + ); + } else { + log::info!( + "[conn={}] Login accepted with authentication disabled", + self.connection_id + ); + } + + if login.packet_size > 0 { + self.packet_size = login.packet_size.min(32767) as u16; + } + if !login.database.is_empty() { + self.database = login.database.clone(); + } + + let session_id = match self.session_pool.checkout(self.db.as_ref()) { + Ok(sid) => sid, + Err(CheckoutError::Exhausted) => { + let err = DbError::Execution("session pool exhausted".to_string()); + let err_resp = build_error_response(&err); + let _ = packet::write_packet(&mut writer, TABULAR_RESULT, &err_resp.data).await; + return Ok(()); + } + }; + self.session_id = Some(session_id); + + if let Err(e) = self.db.set_session_metadata( + session_id, + Some(login.username.clone()), + Some(login.app_name.clone()), + Some(login.hostname.clone()), + Some(self.database.clone()), + ) { + log::error!( + "[conn={}] Failed to set session metadata: {}", + self.connection_id, + e + ); + } + + let mut b = PacketBuilder::with_capacity(512); + tokens::write_envchange_packet_size(&mut b, self.packet_size, self.packet_size); + tokens::write_envchange_database(&mut b, &self.database, "master"); + tokens::write_envchange_language(&mut b, "us_english", ""); + tokens::write_envchange_collation(&mut b); + tokens::write_loginack(&mut b, 0x74000004); + tokens::write_info( + &mut b, + 5701, + 2, + 0, + &format!("Changed database context to '{}'.", &self.database), + "localhost", + "", + 1, + ); + tokens::write_info( + &mut b, + 5703, + 1, + 0, + "Changed language setting to us_english.", + "localhost", + "", + 1, + ); + tokens::write_done(&mut b, tokens::DONE_FINAL, 0, 0); + + if let Err(e) = packet::write_packet(&mut writer, TABULAR_RESULT, b.as_bytes()).await { + self.session_pool.checkin(self.db.as_ref(), session_id); + return Err(format!("Failed to write LOGINACK: {}", e)); + } + log::debug!("[conn={}] Sent LOGINACK response", self.connection_id); + + loop { + let result: Result<_, std::io::Error> = packet::read_message(&mut reader).await; + match result { + Ok((header, data)) => { + crate::session::log_packet(self.connection_id, "packet", &header, &data); + + if header.status & packet::STATUS_RESET != 0 { + log::debug!("[conn={}] STATUS_RESET connection", self.connection_id); + self.prepared_stmts.clear(); + } + + match header.packet_type { + SQL_BATCH => match self.handle_sql_batch(&data, &mut writer).await { + Ok(_) => {} + Err(e) => { + log::error!("[conn={}] SQL batch error: {}", self.connection_id, e); + let err_resp = build_error_response(&e); + let _ = packet::write_packet( + &mut writer, + TABULAR_RESULT, + &err_resp.data, + ) + .await; + } + }, + RPC => match parse_rpc(&data) { + Ok(Some(rpc)) => { + match rpc { + RpcRequest::Sql(sql_req) => { + let preamble = build_param_preamble(&sql_req.params); + let full_sql = if preamble.is_empty() { + sql_req.sql + } else { + format!("{}{}", preamble, sql_req.sql) + }; + match self.execute_sql(full_sql.trim(), &mut writer).await { + Ok(_) => {} + Err(e) => { + log::error!("RPC SQL error: {}", e); + let err_resp = build_error_response(&e); + let _ = packet::write_packet( + &mut writer, + TABULAR_RESULT, + &err_resp.data, + ) + .await; + } + } + } + RpcRequest::Cursor(cursor_req) => { + log::debug!( + "[conn={}] Cursor RPC: {:?}", + self.connection_id, + cursor_req.cursor_op + ); + if let Some(session_id) = self.session_id { + match cursor_req.cursor_op { + CursorOp::Open => { + let sql = cursor_req.sql.unwrap_or_default(); + let scroll_opt = + cursor_req.scroll_opt.unwrap_or(0); + match self.db.cursor_rpc_open( + session_id, &sql, scroll_opt, + ) { + Ok((handle, _result)) => { + let mut buf = PacketBuilder::new(); + tokens::write_output_int( + &mut buf, "@cursor", handle, + ); + tokens::write_done( + &mut buf, + tokens::DONE_FINAL, + 0, + 0, + ); + let _ = packet::write_packet( + &mut writer, + TABULAR_RESULT, + buf.as_bytes(), + ) + .await; + } + Err(e) => { + let err_resp = build_error_response(&e); + let _ = packet::write_packet( + &mut writer, + TABULAR_RESULT, + &err_resp.data, + ) + .await; + } + } + } + CursorOp::Fetch => { + let handle = + cursor_req.cursor_handle.unwrap_or(0); + let fetch_type = + cursor_req.fetch_type.unwrap_or(2); + let row_num = cursor_req.row_num.unwrap_or(0); + let n_rows = cursor_req.n_rows.unwrap_or(1); + match self.db.cursor_rpc_fetch( + session_id, handle, fetch_type, row_num, + n_rows, + ) { + Ok(fetch_result) => { + if fetch_result.rows.is_empty() { + let mut buf = PacketBuilder::new(); + tokens::write_done( + &mut buf, + tokens::DONE_FINAL, + 0, + 0, + ); + let _ = packet::write_packet( + &mut writer, + TABULAR_RESULT, + buf.as_bytes(), + ) + .await; + } else { + let mut buf = PacketBuilder::new(); + let col_types: Vec<_> = fetch_result.column_types.iter() + .map(crate::tds::type_mapping::runtime_type_to_tds) + .collect(); + tokens::write_colmetadata( + &mut buf, + &fetch_result.columns, + &col_types, + ); + for row in &fetch_result.rows { + tokens::write_row( + &mut buf, row, &col_types, + 0, + ); + } + tokens::write_done( + &mut buf, + tokens::DONE_FINAL + | tokens::DONE_COUNT, + 0, + fetch_result.rows.len() as u64, + ); + let _ = packet::write_packet( + &mut writer, + TABULAR_RESULT, + buf.as_bytes(), + ) + .await; + } + } + Err(e) => { + let err_resp = build_error_response(&e); + let _ = packet::write_packet( + &mut writer, + TABULAR_RESULT, + &err_resp.data, + ) + .await; + } + } + } + CursorOp::Close => { + let handle = + cursor_req.cursor_handle.unwrap_or(0); + match self + .db + .cursor_rpc_close(session_id, handle) + { + Ok(()) => { + let mut buf = PacketBuilder::new(); + tokens::write_done( + &mut buf, + tokens::DONE_FINAL, + 0, + 0, + ); + let _ = packet::write_packet( + &mut writer, + TABULAR_RESULT, + buf.as_bytes(), + ) + .await; + } + Err(e) => { + let err_resp = build_error_response(&e); + let _ = packet::write_packet( + &mut writer, + TABULAR_RESULT, + &err_resp.data, + ) + .await; + } + } + } + CursorOp::Unprepare => { + let handle = + cursor_req.cursor_handle.unwrap_or(0); + match self + .db + .cursor_rpc_deallocate(session_id, handle) + { + Ok(()) => { + let mut buf = PacketBuilder::new(); + tokens::write_done( + &mut buf, + tokens::DONE_FINAL, + 0, + 0, + ); + let _ = packet::write_packet( + &mut writer, + TABULAR_RESULT, + buf.as_bytes(), + ) + .await; + } + Err(e) => { + let err_resp = build_error_response(&e); + let _ = packet::write_packet( + &mut writer, + TABULAR_RESULT, + &err_resp.data, + ) + .await; + } + } + } + CursorOp::Prepare + | CursorOp::PrepExec => { + let sql = match cursor_req.sql { + Some(ref s) if !s.is_empty() => s.clone(), + _ => { + let err = DbError::Execution("cursor prepare requires SQL statement".into()); + let err_resp = + build_error_response(&err); + let _ = packet::write_packet( + &mut writer, + TABULAR_RESULT, + &err_resp.data, + ) + .await; + continue; + } + }; + let scroll_opt = + cursor_req.scroll_opt.unwrap_or(0); + match self.db.cursor_rpc_open( + session_id, &sql, scroll_opt, + ) { + Ok((handle, _result)) => { + let mut buf = PacketBuilder::new(); + tokens::write_output_int( + &mut buf, "@cursor", handle, + ); + tokens::write_done( + &mut buf, + tokens::DONE_FINAL, + 0, + 0, + ); + let _ = packet::write_packet( + &mut writer, + TABULAR_RESULT, + buf.as_bytes(), + ) + .await; + } + Err(e) => { + let err_resp = build_error_response(&e); + let _ = packet::write_packet( + &mut writer, + TABULAR_RESULT, + &err_resp.data, + ) + .await; + } + } + } + CursorOp::Execute => { + let handle = + cursor_req.cursor_handle.unwrap_or(0); + let scroll_opt = + cursor_req.scroll_opt.unwrap_or(0); + let n_rows = cursor_req.n_rows.unwrap_or(1); + match self.db.cursor_rpc_fetch( + session_id, handle, 2, 0, n_rows, + ) { + Ok(fetch_result) => { + let _ = scroll_opt; + if fetch_result.rows.is_empty() { + let mut buf = PacketBuilder::new(); + tokens::write_output_int( + &mut buf, "@cursor", handle, + ); + tokens::write_done( + &mut buf, + tokens::DONE_FINAL, + 0, + 0, + ); + let _ = packet::write_packet( + &mut writer, + TABULAR_RESULT, + buf.as_bytes(), + ) + .await; + } else { + let mut buf = PacketBuilder::new(); + let col_types: Vec<_> = fetch_result.column_types.iter() + .map(crate::tds::type_mapping::runtime_type_to_tds) + .collect(); + tokens::write_output_int( + &mut buf, "@cursor", handle, + ); + tokens::write_colmetadata( + &mut buf, + &fetch_result.columns, + &col_types, + ); + for row in &fetch_result.rows { + tokens::write_row( + &mut buf, row, &col_types, + 0, + ); + } + tokens::write_done( + &mut buf, + tokens::DONE_FINAL + | tokens::DONE_COUNT, + 0, + fetch_result.rows.len() as u64, + ); + let _ = packet::write_packet( + &mut writer, + TABULAR_RESULT, + buf.as_bytes(), + ) + .await; + } + } + Err(e) => { + let err_resp = build_error_response(&e); + let _ = packet::write_packet( + &mut writer, + TABULAR_RESULT, + &err_resp.data, + ) + .await; + } + } + } + CursorOp::Option => { + let _ = cursor_req; + let mut buf = PacketBuilder::new(); + tokens::write_done( + &mut buf, + tokens::DONE_FINAL, + 0, + 0, + ); + let _ = packet::write_packet( + &mut writer, + TABULAR_RESULT, + buf.as_bytes(), + ) + .await; + } + } + } else { + let err = DbError::Execution( + "no session for cursor operation".into(), + ); + let err_resp = build_error_response(&err); + let _ = packet::write_packet( + &mut writer, + TABULAR_RESULT, + &err_resp.data, + ) + .await; + } + } + RpcRequest::Prepare(req) => { + let handle = self.prepared_stmts.len() as u32 + 1; + let param_decls = parse_param_decl(&req.param_decl); + let stmt = PreparedStatement { + sql: req.sql.clone(), + param_decls: param_decls.clone(), + }; + self.prepared_stmts.insert(handle, stmt); + let mut buf = PacketBuilder::new(); + tokens::write_output_int( + &mut buf, + "@handle", + handle as i32, + ); + tokens::write_done(&mut buf, tokens::DONE_FINAL, 0, 0); + let _ = packet::write_packet( + &mut writer, + TABULAR_RESULT, + buf.as_bytes(), + ) + .await; + } + RpcRequest::Execute(req) => { + let handle = req.stmt_handle as u32; + if let Some(stmt) = self.prepared_stmts.get(&handle) { + let param_decls = &stmt.param_decls; + let preamble = build_param_preamble_with_decls( + &req.params, + param_decls, + ); + let full_sql = if preamble.is_empty() { + stmt.sql.clone() + } else { + format!("{}{}", preamble, stmt.sql) + }; + match self + .execute_sql(full_sql.trim(), &mut writer) + .await + { + Ok(_) => {} + Err(e) => { + log::error!("sp_execute error: {}", e); + let err_resp = build_error_response(&e); + let _ = packet::write_packet( + &mut writer, + TABULAR_RESULT, + &err_resp.data, + ) + .await; + } + } + } else { + let err = DbError::Execution(format!( + "Invalid statement handle: {}", + req.stmt_handle + )); + let err_resp = build_error_response(&err); + let _ = packet::write_packet( + &mut writer, + TABULAR_RESULT, + &err_resp.data, + ) + .await; + } + } + RpcRequest::Unprepare(req) => { + let handle = req.stmt_handle as u32; + self.prepared_stmts.remove(&handle); + let mut buf = PacketBuilder::new(); + tokens::write_done(&mut buf, tokens::DONE_FINAL, 0, 0); + let _ = packet::write_packet( + &mut writer, + TABULAR_RESULT, + buf.as_bytes(), + ) + .await; + } + RpcRequest::PrepExec(req) => { + let handle = self.prepared_stmts.len() as u32 + 1; + let param_decls = parse_param_decl(&req.param_decl); + let stmt = PreparedStatement { + sql: req.sql.clone(), + param_decls: param_decls.clone(), + }; + self.prepared_stmts.insert(handle, stmt); + let preamble = build_param_preamble_with_decls( + &req.params, + ¶m_decls, + ); + let full_sql = if preamble.is_empty() { + req.sql.clone() + } else { + format!("{}{}", preamble, req.sql) + }; + match self.execute_sql(full_sql.trim(), &mut writer).await { + Ok(_) => {} + Err(e) => { + log::error!("sp_prepexec error: {}", e); + let err_resp = build_error_response(&e); + let _ = packet::write_packet( + &mut writer, + TABULAR_RESULT, + &err_resp.data, + ) + .await; + } + } + } + RpcRequest::ResetConnection => { + self.prepared_stmts.clear(); + let mut buf = PacketBuilder::new(); + tokens::write_done(&mut buf, tokens::DONE_FINAL, 0, 0); + let _ = packet::write_packet( + &mut writer, + TABULAR_RESULT, + buf.as_bytes(), + ) + .await; + } + RpcRequest::Catalog(cat_req) => { + let sql = match cat_req.proc { + CatalogProc::Tables => { + let table_name = cat_req + .params.first() + .map(|p| p.value_sql.trim_matches('\'')) + .unwrap_or("%"); + let table_owner = cat_req + .params + .get(1) + .map(|p| p.value_sql.trim_matches('\'')) + .unwrap_or("%"); + format!( + "SELECT DB_NAME() AS TABLE_QUALIFIER, s.name AS TABLE_OWNER, t.name AS TABLE_NAME, 'TABLE' AS TABLE_TYPE, NULL AS REMARKS FROM sys.tables t JOIN sys.schemas s ON t.schema_id = s.id WHERE t.name LIKE '{}' AND s.name LIKE '{}'", + table_name, table_owner + ) + } + CatalogProc::Columns => { + let table_name = cat_req + .params.first() + .map(|p| p.value_sql.trim_matches('\'')) + .unwrap_or("%"); + format!( + "SELECT s.name AS TABLE_OWNER, t.name AS TABLE_NAME, c.name AS COLUMN_NAME, c.column_id AS ORDINAL_POSITION, ty.name AS TYPE_NAME FROM sys.columns c JOIN sys.tables t ON c.object_id = t.object_id JOIN sys.schemas s ON t.schema_id = s.id JOIN sys.types ty ON c.user_type_id = ty.user_type_id WHERE t.name LIKE '{}' ORDER BY t.name, c.column_id", + table_name + ) + } + CatalogProc::SprocColumns => { + let proc_name = cat_req + .params.first() + .map(|p| p.value_sql.trim_matches('\'')) + .unwrap_or("%"); + format!( + "SELECT s.name AS PROCEDURE_OWNER, r.name AS PROCEDURE_NAME, p.name AS COLUMN_NAME, p.parameter_id AS ORDINAL_POSITION, ty.name AS TYPE_NAME FROM sys.parameters p JOIN sys.routines r ON p.object_id = r.object_id JOIN sys.schemas s ON r.schema_id = s.id JOIN sys.types ty ON p.user_type_id = ty.user_type_id WHERE r.name LIKE '{}' ORDER BY r.name, p.parameter_id", + proc_name + ) + } + CatalogProc::PrimaryKeys => { + let table_name = cat_req + .params.first() + .map(|p| p.value_sql.trim_matches('\'')) + .unwrap_or("%"); + format!( + "SELECT s.name AS TABLE_OWNER, t.name AS TABLE_NAME, c.name AS COLUMN_NAME, c.column_id AS KEY_SEQ, pk.name AS PK_NAME FROM sys.columns c JOIN sys.tables t ON c.object_id = t.object_id JOIN sys.schemas s ON t.schema_id = s.id JOIN (SELECT ic.object_id, ic.column_id, i.name FROM sys.index_columns ic JOIN sys.indexes i ON ic.object_id = i.object_id AND ic.index_id = i.index_id WHERE i.is_primary_key = 1) pk ON c.object_id = pk.object_id AND c.column_id = pk.column_id WHERE t.name LIKE '{}' ORDER BY t.name, c.column_id", + table_name + ) + } + CatalogProc::DescribeCursor => { + let cursor_handle = cat_req + .params + .get(2) + .and_then(|p| p.value_sql.parse::().ok()) + .unwrap_or(0); + if let Some(session_id) = self.session_id { + match self.db.cursor_rpc_fetch( + session_id, + cursor_handle, + 2, + 0, + 1, + ) { + Ok(fetch_result) => { + let mut buf = PacketBuilder::new(); + let col_names = vec![ + "reference_name".to_string(), + "cursor_name".to_string(), + "cursor_scope".to_string(), + "status".to_string(), + "model".to_string(), + "concurrency".to_string(), + "scrollable".to_string(), + "open_status".to_string(), + "cursor_rows".to_string(), + "fetch_status".to_string(), + "column_count".to_string(), + "row_count".to_string(), + "last_operation".to_string(), + "cursor_handle".to_string(), + ]; + let col_types: Vec<_> = col_names.iter() + .map(|_| crate::tds::type_mapping::runtime_type_to_tds(&iridium_core::types::DataType::Int)) + .collect(); + let cursor_name = format!( + "#rpc_cursor_{}", + cursor_handle + ); + let row = vec![ + iridium_core::types::Value::NVarChar(cursor_name.clone()), + iridium_core::types::Value::NVarChar(cursor_name), + iridium_core::types::Value::Int(1), + iridium_core::types::Value::Int(0), + iridium_core::types::Value::Int(0), + iridium_core::types::Value::Int(1), + iridium_core::types::Value::Int(1), + iridium_core::types::Value::Int(if fetch_result.rows.is_empty() { 0 } else { 1 }), + iridium_core::types::Value::Int(fetch_result.rows.len() as i32), + iridium_core::types::Value::Int(fetch_result.fetch_status), + iridium_core::types::Value::Int(fetch_result.columns.len() as i32), + iridium_core::types::Value::Int(fetch_result.rows.len() as i32), + iridium_core::types::Value::Int(0), + iridium_core::types::Value::Int(cursor_handle), + ]; + tokens::write_colmetadata( + &mut buf, &col_names, &col_types, + ); + tokens::write_row( + &mut buf, &row, &col_types, 0, + ); + tokens::write_done( + &mut buf, + tokens::DONE_FINAL + | tokens::DONE_COUNT, + 0, + 1, + ); + let _ = packet::write_packet( + &mut writer, + TABULAR_RESULT, + buf.as_bytes(), + ) + .await; + } + Err(e) => { + let err_resp = build_error_response(&e); + let _ = packet::write_packet( + &mut writer, + TABULAR_RESULT, + &err_resp.data, + ) + .await; + } + } + } else { + let err = DbError::Execution( + "no session for cursor operation".into(), + ); + let err_resp = build_error_response(&err); + let _ = packet::write_packet( + &mut writer, + TABULAR_RESULT, + &err_resp.data, + ) + .await; + } + String::new() + } + }; + if !sql.is_empty() { + match self.execute_sql(&sql, &mut writer).await { + Ok(_) => {} + Err(e) => { + log::error!("Catalog RPC error: {}", e); + let err_resp = build_error_response(&e); + let _ = packet::write_packet( + &mut writer, + TABULAR_RESULT, + &err_resp.data, + ) + .await; + } + } + } + } + } + } + Ok(None) => { + let preview_len = data.len().min(96); + log::warn!( + "[conn={}] Unsupported RPC request ({} bytes), first {} bytes: {:02X?}", + self.connection_id, + data.len(), + preview_len, + &data[..preview_len] + ); + let err = DbError::Parse("unsupported RPC request".into()); + let err_resp = build_error_response(&err); + let _ = packet::write_packet( + &mut writer, + TABULAR_RESULT, + &err_resp.data, + ) + .await; + } + Err(e) => { + let err = iridium_core::error::DbError::Parse(e.to_string()); + let err_resp = build_error_response(&err); + let _ = packet::write_packet( + &mut writer, + TABULAR_RESULT, + &err_resp.data, + ) + .await; + } + }, + BULK_LOAD => { + log::info!("[conn={}] BULK_LOAD received", self.connection_id); + if let Some(session_id) = self.session_id { + let (active, target, columns, received_metadata) = self + .db + .session_options(session_id) + .map(|_opts| self.db.get_bulk_load_state(session_id)) + .unwrap_or((false, None, None, false)); + + if active && target.is_some() && columns.is_some() { + let target = target.unwrap(); + let columns = columns.unwrap(); + + let mut reader = packet::PacketReader::new(&data); + let mut column_types = Vec::new(); + + if !received_metadata { + let token = reader.read_u8().map_err(|e| e.to_string())?; + if token != tokens::COLMETADATA_TOKEN { + return Err(format!( + "Expected COLMETADATA (0x81), got 0x{:02X}", + token + )); + } + let count = + reader.read_u16_le().map_err(|e| e.to_string())? + as usize; + for _ in 0..count { + reader.skip(4).map_err(|e| e.to_string())?; + let _flags = + reader.read_u16_le().map_err(|e| e.to_string())?; + let ti = crate::tds::type_mapping::read_type_info( + &mut reader, + ) + .map_err(|e| e.to_string())?; + column_types.push(ti); + let name_len = + reader.read_u8().map_err(|e| e.to_string())? + as usize; + let _name = reader + .read_utf16le(name_len) + .map_err(|e| e.to_string())?; + } + self.db + .set_bulk_load_active( + session_id, + true, + target.clone(), + columns.clone(), + true, + ) + .map_err(|e| e.to_string())?; + } + + match parse_bulk_load_data(&data, &columns) { + Ok(bulk_data) => { + log::info!( + "[conn={}] Parsed {} bulk rows for table {}", + self.connection_id, + bulk_data.rows.len(), + target.name + ); + + let mut sql = String::new(); + let col_names: Vec = + columns.iter().map(|c| c.name.clone()).collect(); + let col_list = col_names.join(", "); + + for row in bulk_data.rows { + let vals: Vec = row + .iter() + .map(|v| v.to_sql_literal()) + .collect(); + sql.push_str(&format!( + "INSERT INTO {}.{} ({}) VALUES ({});\n", + target.schema_or_dbo(), + target.name, + col_list, + vals.join(", ") + )); + } + + self.db + .set_bulk_load_active( + session_id, + false, + target.clone(), + columns.clone(), + false, + ) + .map_err(|e| e.to_string())?; + + match self.execute_sql(&sql, &mut writer).await { + Ok(_) => {} + Err(e) => { + log::error!( + "Bulk insert execution error: {}", + e + ); + let err_resp = build_error_response(&e); + let _ = packet::write_packet( + &mut writer, + TABULAR_RESULT, + &err_resp.data, + ) + .await; + } + } + } + Err(e) => { + log::error!("Bulk data parse error: {}", e); + let err = DbError::Parse(e.to_string()); + let err_resp = build_error_response(&err); + let _ = packet::write_packet( + &mut writer, + TABULAR_RESULT, + &err_resp.data, + ) + .await; + } + } + } else { + log::warn!("[conn={}] Received BULK_LOAD but bulk load was not expected", self.connection_id); + } + } + } + ATTENTION => { + log::debug!("[conn={}] ATTENTION received", self.connection_id); + let mut attn = PacketBuilder::new(); + tokens::write_done(&mut attn, tokens::DONE_ATTN, 0, 0); + let _ = + packet::write_packet(&mut writer, TABULAR_RESULT, attn.as_bytes()) + .await; + } + _ => { + log::warn!( + "[conn={}] Unsupported packet type 0x{:02X}", + self.connection_id, + header.packet_type + ); + } + } + } + Err(ref e) if e.kind() == std::io::ErrorKind::UnexpectedEof => { + log::debug!("[conn={}] Client disconnected", self.connection_id); + break; + } + Err(e) => { + log::error!("[conn={}] Read error: {}", self.connection_id, e); + break; + } + } + } + + if let Some(sid) = self.session_id.take() { + self.session_pool.checkin(self.db.as_ref(), sid); + } + + Ok(()) + } +} diff --git a/crates/iridium_server/src/session/mod.rs b/crates/iridium_server/src/session/mod.rs index 03bfad4..fe276a5 100644 --- a/crates/iridium_server/src/session/mod.rs +++ b/crates/iridium_server/src/session/mod.rs @@ -1,44 +1,23 @@ use once_cell::sync::Lazy; use regex::Regex; use std::collections::HashMap; -use std::sync::atomic::{AtomicBool, AtomicU64, Ordering}; +use std::sync::atomic::{AtomicU64, Ordering}; use std::sync::Arc; -use tokio::io::AsyncWriteExt; -use tokio::net::TcpStream; -use iridium_core::types::{DataType, Value}; -use iridium_core::{error::DbError, SessionId}; +use iridium_core::SessionId; -use super::pool::{CheckoutError, SessionPool}; +use super::pool::SessionPool; use crate::ServerDatabase; -mod compat; -mod response; - -use self::compat::{ - extract_leading_use_database, is_ssms_contained_auth_probe, parse_simple_use_database, -}; -use self::response::{build_single_int_result, build_use_database_response}; -use super::tds::batch::{build_error_response, parse_sql_batch}; -use super::tds::bulk::parse_bulk_load_data; -use super::tds::login::parse_login7; -use super::tds::packet::{ - self, PacketBuilder, ATTENTION, BULK_LOAD, RPC, SQL_BATCH, TABULAR_RESULT, TDS7_LOGIN, - TDS7_PRELOGIN, -}; -use super::tds::prelogin::{ - build_prelogin_response, parse_prelogin, ENCRYPT_NOT_SUP, ENCRYPT_OFF, ENCRYPT_ON, - ENCRYPT_REQUIRED, -}; -use super::tds::rpc::{ - build_param_preamble, build_param_preamble_with_decls, parse_param_decl, parse_rpc, - CatalogProc, RpcRequest, -}; -use super::tds::tokens; -use super::tds_tls_io::TdsTlsIo; -use super::tls; use super::ServerConfig; -trait AsyncReadWrite: tokio::io::AsyncRead + tokio::io::AsyncWrite + Send + Unpin {} +pub mod compat; +pub mod response; +pub mod execution; +pub mod handshake; + +use crate::tds::packet; + +pub(crate) trait AsyncReadWrite: tokio::io::AsyncRead + tokio::io::AsyncWrite + Send + Unpin {} impl AsyncReadWrite for T {} static NEXT_CONNECTION_ID: AtomicU64 = AtomicU64::new(1); @@ -78,1362 +57,6 @@ impl TdsSession { prepared_stmts: HashMap::new(), } } - - pub async fn handle(self, mut stream: TcpStream) -> Result<(), String> { - let mut needs_tls_upgrade = false; - let mut login_packet = None; - - log::info!( - "[conn={}] Starting handshake for incoming connection", - self.connection_id - ); - loop { - let (header, data) = packet::read_message(&mut stream) - .await - .map_err(|e| format!("Handshake read error: {}", e))?; - log_packet(self.connection_id, "handshake", &header, &data); - - if header.packet_type == TDS7_PRELOGIN { - log::debug!( - "[conn={}] PRELOGIN data hex: {:02X?}", - self.connection_id, - data - ); - let prelogin = parse_prelogin(&data).map_err(|e| e.to_string())?; - log::debug!( - "[conn={}] PRELOGIN: version={:?}, encryption={}", - self.connection_id, - prelogin.version, - prelogin.encryption - ); - - // When TLS is disabled, respond based on client's request: - // - If client requested ENCRYPT_NOT_SUP, echo ENCRYPT_NOT_SUP - // - If client requested ENCRYPT_ON/REQUIRED, respond with ENCRYPT_NOT_SUP - // - If client requested ENCRYPT_OFF, respond with ENCRYPT_OFF - // This maintains compatibility with clients that use strict - // prelogin encryption negotiation. - let server_encrypt = if self.config.tls_enabled { - ENCRYPT_ON - } else if prelogin.encryption == ENCRYPT_NOT_SUP - || prelogin.encryption == ENCRYPT_ON - || prelogin.encryption == ENCRYPT_REQUIRED - { - ENCRYPT_NOT_SUP - } else { - ENCRYPT_OFF - }; - - needs_tls_upgrade = self.config.tls_enabled - && (prelogin.encryption == ENCRYPT_ON - || prelogin.encryption == ENCRYPT_REQUIRED); - - let response = build_prelogin_response(server_encrypt); - packet::write_packet(&mut stream, TDS7_PRELOGIN, &response) - .await - .map_err(|e| format!("Failed to write PRELOGIN response: {}", e))?; - stream - .flush() - .await - .map_err(|e| format!("Failed to flush: {}", e))?; - log::debug!( - "[conn={}] Sent PRELOGIN response (encryption={})", - self.connection_id, - server_encrypt - ); - - if needs_tls_upgrade { - break; - } - } else if header.packet_type == TDS7_LOGIN { - log_packet(self.connection_id, "login", &header, &data); - login_packet = Some((header, data)); - break; - } else { - return Err(format!( - "Expected PRELOGIN or LOGIN7, got 0x{:02X}", - header.packet_type - )); - } - } - - self.handle_login(stream, login_packet, needs_tls_upgrade) - .await - } - - async fn handle_login( - mut self, - stream: TcpStream, - login_packet: Option<(packet::PacketHeader, Vec)>, - needs_tls_upgrade: bool, - ) -> Result<(), String> { - // Perform TLS upgrade if needed. - // In TDS 7.4, the TLS handshake is tunneled inside TDS PRELOGIN (0x12) - // packets. After handshake completion, traffic switches to raw TLS over TCP. - let stream: Box = if needs_tls_upgrade { - log::info!( - "[conn={}] Client requested TLS, upgrading connection via TDS-wrapped handshake", - self.connection_id - ); - - let tls_config = if let (Some(cert_path), Some(key_path)) = - (&self.config.tls_cert_path, &self.config.tls_key_path) - { - tls::load_tls_config(cert_path, key_path) - .map_err(|e| format!("Failed to load TLS config: {}", e))? - } else { - return Err("TLS enabled but no certificate configured".to_string()); - }; - - let acceptor = tokio_rustls::TlsAcceptor::from(std::sync::Arc::new(tls_config)); - - // Use TdsTlsIo to strip/wrap TDS framing during handshake - let raw_mode = Arc::new(AtomicBool::new(false)); - let tds_io = TdsTlsIo::new(stream, raw_mode.clone()); - - let tls_stream = acceptor - .accept(tds_io) - .await - .map_err(|e| format!("TLS handshake failed: {}", e))?; - - // Switch to raw mode — post-handshake traffic is raw TLS over TCP - raw_mode.store(true, Ordering::Release); - - log::info!( - "[conn={}] TLS handshake completed, switched to raw mode", - self.connection_id - ); - Box::new(tls_stream) - } else { - Box::new(stream) - }; - - // Now use the stream for the rest - let (mut reader, mut writer) = tokio::io::split(stream); - - let login_data = if let Some((_, data)) = login_packet { - data - } else { - let (header, data) = packet::read_message(&mut reader) - .await - .map_err(|e| format!("Failed to read LOGIN7: {}", e))?; - if header.packet_type != TDS7_LOGIN { - return Err(format!( - "Expected LOGIN7 (0x10), got 0x{:02X}", - header.packet_type - )); - } - data - }; - - let login = parse_login7(&login_data).map_err(|e| e.to_string())?; - log::info!( - "[conn={}] LOGIN7: user={}, database={}, app={}, packet_size={}", - self.connection_id, - login.username, - login.database, - login.app_name, - login.packet_size - ); - - if let Some(ref creds) = self.config.auth { - if login.username != creds.user || login.password != creds.password { - log::warn!( - "[conn={}] Login rejected for user={} against configured SQL auth", - self.connection_id, - login.username - ); - let err = DbError::Execution("Login failed for user.".to_string()); - let err_resp = build_error_response(&err); - packet::write_packet(&mut writer, TABULAR_RESULT, &err_resp.data) - .await - .map_err(|e| e.to_string())?; - return Ok(()); - } - log::info!( - "[conn={}] Login accepted for user={}", - self.connection_id, - login.username - ); - } else { - log::info!( - "[conn={}] Login accepted with authentication disabled", - self.connection_id - ); - } - - if login.packet_size > 0 { - self.packet_size = login.packet_size.min(32767) as u16; - } - if !login.database.is_empty() { - self.database = login.database.clone(); - } - - let session_id = match self.session_pool.checkout(self.db.as_ref()) { - Ok(sid) => sid, - Err(CheckoutError::Exhausted) => { - let err = DbError::Execution("session pool exhausted".to_string()); - let err_resp = build_error_response(&err); - packet::write_packet(&mut writer, TABULAR_RESULT, &err_resp.data) - .await - .map_err(|e| e.to_string())?; - return Ok(()); - } - }; - self.session_id = Some(session_id); - - // Set session metadata - if let Err(e) = self.db.set_session_metadata( - session_id, - Some(login.username.clone()), - Some(login.app_name.clone()), - Some(login.hostname.clone()), - Some(self.database.clone()), - ) { - log::error!( - "[conn={}] Failed to set session metadata: {}", - self.connection_id, - e - ); - } - - // Build LOGINACK response - // SSMS expects the full sequence: ENVCHANGE(PacketSize), ENVCHANGE(Database), - // ENVCHANGE(Language), ENVCHANGE(Collation), LOGINACK, INFO, DONE_FINAL - let mut b = PacketBuilder::with_capacity(512); - tokens::write_envchange_packet_size(&mut b, self.packet_size, self.packet_size); - tokens::write_envchange_database(&mut b, &self.database, "master"); - tokens::write_envchange_language(&mut b, "us_english", ""); - tokens::write_envchange_collation(&mut b); - tokens::write_loginack(&mut b, 0x74000004); - tokens::write_info( - &mut b, - 5701, - 2, - 0, - &format!("Changed database context to '{}'.", &self.database), - "localhost", - "", - 1, - ); - tokens::write_info( - &mut b, - 5703, - 1, - 0, - "Changed language setting to us_english.", - "localhost", - "", - 1, - ); - tokens::write_done(&mut b, tokens::DONE_FINAL, 0, 0); - - if let Err(e) = packet::write_packet(&mut writer, TABULAR_RESULT, b.as_bytes()).await { - self.session_pool.checkin(self.db.as_ref(), session_id); - return Err(format!("Failed to write LOGINACK: {}", e)); - } - log::debug!("[conn={}] Sent LOGINACK response", self.connection_id); - - // Main loop: handle SQL Batch and other packets - loop { - let result = packet::read_message(&mut reader).await; - match result { - Ok((header, data)) => { - log_packet(self.connection_id, "packet", &header, &data); - - // Handle STATUS_RESET (0x08) in TDS packet header - reset connection state before processing - if header.status & packet::STATUS_RESET != 0 { - log::debug!("[conn={}] STATUS_RESET connection", self.connection_id); - self.prepared_stmts.clear(); - } - - match header.packet_type { - SQL_BATCH => match self.handle_sql_batch(&data, &mut writer).await { - Ok(_) => {} - Err(e) => { - log::error!("[conn={}] SQL batch error: {}", self.connection_id, e); - let err_resp = build_error_response(&e); - let _ = packet::write_packet( - &mut writer, - TABULAR_RESULT, - &err_resp.data, - ) - .await; - } - }, - RPC => match parse_rpc(&data) { - Ok(Some(rpc)) => { - match rpc { - RpcRequest::Sql(sql_req) => { - let preamble = build_param_preamble(&sql_req.params); - let full_sql = if preamble.is_empty() { - sql_req.sql - } else { - format!("{}{}", preamble, sql_req.sql) - }; - match self.execute_sql(full_sql.trim(), &mut writer).await { - Ok(_) => {} - Err(e) => { - log::error!("RPC SQL error: {}", e); - let err_resp = build_error_response(&e); - let _ = packet::write_packet( - &mut writer, - TABULAR_RESULT, - &err_resp.data, - ) - .await; - } - } - } - RpcRequest::Cursor(cursor_req) => { - log::debug!( - "[conn={}] Cursor RPC: {:?}", - self.connection_id, - cursor_req.cursor_op - ); - if let Some(session_id) = self.session_id { - match cursor_req.cursor_op { - super::tds::rpc::CursorOp::Open => { - let sql = cursor_req.sql.unwrap_or_default(); - let scroll_opt = - cursor_req.scroll_opt.unwrap_or(0); - match self.db.cursor_rpc_open( - session_id, &sql, scroll_opt, - ) { - Ok((handle, _result)) => { - let mut buf = PacketBuilder::new(); - tokens::write_output_int( - &mut buf, "@cursor", handle, - ); - tokens::write_done( - &mut buf, - tokens::DONE_FINAL, - 0, - 0, - ); - let _ = packet::write_packet( - &mut writer, - TABULAR_RESULT, - buf.as_bytes(), - ) - .await; - } - Err(e) => { - let err_resp = build_error_response(&e); - let _ = packet::write_packet( - &mut writer, - TABULAR_RESULT, - &err_resp.data, - ) - .await; - } - } - } - super::tds::rpc::CursorOp::Fetch => { - let handle = - cursor_req.cursor_handle.unwrap_or(0); - let fetch_type = - cursor_req.fetch_type.unwrap_or(2); // Default NEXT - let row_num = cursor_req.row_num.unwrap_or(0); - let n_rows = cursor_req.n_rows.unwrap_or(1); - match self.db.cursor_rpc_fetch( - session_id, handle, fetch_type, row_num, - n_rows, - ) { - Ok(fetch_result) => { - if fetch_result.rows.is_empty() { - let mut buf = PacketBuilder::new(); - tokens::write_done( - &mut buf, - tokens::DONE_FINAL, - 0, - 0, - ); - let _ = packet::write_packet( - &mut writer, - TABULAR_RESULT, - buf.as_bytes(), - ) - .await; - } else { - let mut buf = PacketBuilder::new(); - let col_types: Vec<_> = fetch_result.column_types.iter() - .map(super::tds::type_mapping::runtime_type_to_tds) - .collect(); - tokens::write_colmetadata( - &mut buf, - &fetch_result.columns, - &col_types, - ); - for row in &fetch_result.rows { - tokens::write_row( - &mut buf, row, &col_types, - 0, - ); - } - tokens::write_done( - &mut buf, - tokens::DONE_FINAL - | tokens::DONE_COUNT, - 0, - fetch_result.rows.len() as u64, - ); - let _ = packet::write_packet( - &mut writer, - TABULAR_RESULT, - buf.as_bytes(), - ) - .await; - } - } - Err(e) => { - let err_resp = build_error_response(&e); - let _ = packet::write_packet( - &mut writer, - TABULAR_RESULT, - &err_resp.data, - ) - .await; - } - } - } - super::tds::rpc::CursorOp::Close => { - let handle = - cursor_req.cursor_handle.unwrap_or(0); - match self - .db - .cursor_rpc_close(session_id, handle) - { - Ok(()) => { - let mut buf = PacketBuilder::new(); - tokens::write_done( - &mut buf, - tokens::DONE_FINAL, - 0, - 0, - ); - let _ = packet::write_packet( - &mut writer, - TABULAR_RESULT, - buf.as_bytes(), - ) - .await; - } - Err(e) => { - let err_resp = build_error_response(&e); - let _ = packet::write_packet( - &mut writer, - TABULAR_RESULT, - &err_resp.data, - ) - .await; - } - } - } - super::tds::rpc::CursorOp::Unprepare => { - let handle = - cursor_req.cursor_handle.unwrap_or(0); - match self - .db - .cursor_rpc_deallocate(session_id, handle) - { - Ok(()) => { - let mut buf = PacketBuilder::new(); - tokens::write_done( - &mut buf, - tokens::DONE_FINAL, - 0, - 0, - ); - let _ = packet::write_packet( - &mut writer, - TABULAR_RESULT, - buf.as_bytes(), - ) - .await; - } - Err(e) => { - let err_resp = build_error_response(&e); - let _ = packet::write_packet( - &mut writer, - TABULAR_RESULT, - &err_resp.data, - ) - .await; - } - } - } - super::tds::rpc::CursorOp::Prepare - | super::tds::rpc::CursorOp::PrepExec => { - let sql = match cursor_req.sql { - Some(ref s) if !s.is_empty() => s.clone(), - _ => { - let err = DbError::Execution("cursor prepare requires SQL statement".into()); - let err_resp = - build_error_response(&err); - let _ = packet::write_packet( - &mut writer, - TABULAR_RESULT, - &err_resp.data, - ) - .await; - continue; - } - }; - let scroll_opt = - cursor_req.scroll_opt.unwrap_or(0); - match self.db.cursor_rpc_open( - session_id, &sql, scroll_opt, - ) { - Ok((handle, _result)) => { - let mut buf = PacketBuilder::new(); - tokens::write_output_int( - &mut buf, "@cursor", handle, - ); - tokens::write_done( - &mut buf, - tokens::DONE_FINAL, - 0, - 0, - ); - let _ = packet::write_packet( - &mut writer, - TABULAR_RESULT, - buf.as_bytes(), - ) - .await; - } - Err(e) => { - let err_resp = build_error_response(&e); - let _ = packet::write_packet( - &mut writer, - TABULAR_RESULT, - &err_resp.data, - ) - .await; - } - } - } - super::tds::rpc::CursorOp::Execute => { - let handle = - cursor_req.cursor_handle.unwrap_or(0); - let scroll_opt = - cursor_req.scroll_opt.unwrap_or(0); - let n_rows = cursor_req.n_rows.unwrap_or(1); - match self.db.cursor_rpc_fetch( - session_id, handle, 2, 0, n_rows, - ) { - Ok(fetch_result) => { - let _ = scroll_opt; - if fetch_result.rows.is_empty() { - let mut buf = PacketBuilder::new(); - tokens::write_output_int( - &mut buf, "@cursor", handle, - ); - tokens::write_done( - &mut buf, - tokens::DONE_FINAL, - 0, - 0, - ); - let _ = packet::write_packet( - &mut writer, - TABULAR_RESULT, - buf.as_bytes(), - ) - .await; - } else { - let mut buf = PacketBuilder::new(); - let col_types: Vec<_> = fetch_result.column_types.iter() - .map(super::tds::type_mapping::runtime_type_to_tds) - .collect(); - tokens::write_output_int( - &mut buf, "@cursor", handle, - ); - tokens::write_colmetadata( - &mut buf, - &fetch_result.columns, - &col_types, - ); - for row in &fetch_result.rows { - tokens::write_row( - &mut buf, row, &col_types, - 0, - ); - } - tokens::write_done( - &mut buf, - tokens::DONE_FINAL - | tokens::DONE_COUNT, - 0, - fetch_result.rows.len() as u64, - ); - let _ = packet::write_packet( - &mut writer, - TABULAR_RESULT, - buf.as_bytes(), - ) - .await; - } - } - Err(e) => { - let err_resp = build_error_response(&e); - let _ = packet::write_packet( - &mut writer, - TABULAR_RESULT, - &err_resp.data, - ) - .await; - } - } - } - super::tds::rpc::CursorOp::Option => { - let _ = cursor_req; - let mut buf = PacketBuilder::new(); - tokens::write_done( - &mut buf, - tokens::DONE_FINAL, - 0, - 0, - ); - let _ = packet::write_packet( - &mut writer, - TABULAR_RESULT, - buf.as_bytes(), - ) - .await; - } - } - } else { - let err = DbError::Execution( - "no session for cursor operation".into(), - ); - let err_resp = build_error_response(&err); - let _ = packet::write_packet( - &mut writer, - TABULAR_RESULT, - &err_resp.data, - ) - .await; - } - } - RpcRequest::Prepare(req) => { - let handle = self.prepared_stmts.len() as u32 + 1; - let param_decls = parse_param_decl(&req.param_decl); - let stmt = PreparedStatement { - sql: req.sql.clone(), - param_decls: param_decls.clone(), - }; - self.prepared_stmts.insert(handle, stmt); - let mut buf = PacketBuilder::new(); - tokens::write_output_int( - &mut buf, - "@handle", - handle as i32, - ); - tokens::write_done(&mut buf, tokens::DONE_FINAL, 0, 0); - let _ = packet::write_packet( - &mut writer, - TABULAR_RESULT, - buf.as_bytes(), - ) - .await; - } - RpcRequest::Execute(req) => { - let handle = req.stmt_handle as u32; - if let Some(stmt) = self.prepared_stmts.get(&handle) { - let param_decls = &stmt.param_decls; - let preamble = build_param_preamble_with_decls( - &req.params, - param_decls, - ); - let full_sql = if preamble.is_empty() { - stmt.sql.clone() - } else { - format!("{}{}", preamble, stmt.sql) - }; - match self - .execute_sql(full_sql.trim(), &mut writer) - .await - { - Ok(_) => {} - Err(e) => { - log::error!("sp_execute error: {}", e); - let err_resp = build_error_response(&e); - let _ = packet::write_packet( - &mut writer, - TABULAR_RESULT, - &err_resp.data, - ) - .await; - } - } - } else { - let err = DbError::Execution(format!( - "Invalid statement handle: {}", - req.stmt_handle - )); - let err_resp = build_error_response(&err); - let _ = packet::write_packet( - &mut writer, - TABULAR_RESULT, - &err_resp.data, - ) - .await; - } - } - RpcRequest::Unprepare(req) => { - let handle = req.stmt_handle as u32; - self.prepared_stmts.remove(&handle); - let mut buf = PacketBuilder::new(); - tokens::write_done(&mut buf, tokens::DONE_FINAL, 0, 0); - let _ = packet::write_packet( - &mut writer, - TABULAR_RESULT, - buf.as_bytes(), - ) - .await; - } - RpcRequest::PrepExec(req) => { - let handle = self.prepared_stmts.len() as u32 + 1; - let param_decls = parse_param_decl(&req.param_decl); - let stmt = PreparedStatement { - sql: req.sql.clone(), - param_decls: param_decls.clone(), - }; - self.prepared_stmts.insert(handle, stmt); - let preamble = build_param_preamble_with_decls( - &req.params, - ¶m_decls, - ); - let full_sql = if preamble.is_empty() { - req.sql.clone() - } else { - format!("{}{}", preamble, req.sql) - }; - match self.execute_sql(full_sql.trim(), &mut writer).await { - Ok(_) => {} - Err(e) => { - log::error!("sp_prepexec error: {}", e); - let err_resp = build_error_response(&e); - let _ = packet::write_packet( - &mut writer, - TABULAR_RESULT, - &err_resp.data, - ) - .await; - } - } - } - RpcRequest::ResetConnection => { - self.prepared_stmts.clear(); - let mut buf = PacketBuilder::new(); - tokens::write_done(&mut buf, tokens::DONE_FINAL, 0, 0); - let _ = packet::write_packet( - &mut writer, - TABULAR_RESULT, - buf.as_bytes(), - ) - .await; - } - RpcRequest::Catalog(cat_req) => { - let sql = match cat_req.proc { - CatalogProc::Tables => { - let table_name = cat_req - .params.first() - .map(|p| p.value_sql.trim_matches('\'')) - .unwrap_or("%"); - let table_owner = cat_req - .params - .get(1) - .map(|p| p.value_sql.trim_matches('\'')) - .unwrap_or("%"); - format!( - "SELECT DB_NAME() AS TABLE_QUALIFIER, s.name AS TABLE_OWNER, t.name AS TABLE_NAME, 'TABLE' AS TABLE_TYPE, NULL AS REMARKS FROM sys.tables t JOIN sys.schemas s ON t.schema_id = s.id WHERE t.name LIKE '{}' AND s.name LIKE '{}'", - table_name, table_owner - ) - } - CatalogProc::Columns => { - let table_name = cat_req - .params.first() - .map(|p| p.value_sql.trim_matches('\'')) - .unwrap_or("%"); - format!( - "SELECT s.name AS TABLE_OWNER, t.name AS TABLE_NAME, c.name AS COLUMN_NAME, c.column_id AS ORDINAL_POSITION, ty.name AS TYPE_NAME FROM sys.columns c JOIN sys.tables t ON c.object_id = t.object_id JOIN sys.schemas s ON t.schema_id = s.id JOIN sys.types ty ON c.user_type_id = ty.user_type_id WHERE t.name LIKE '{}' ORDER BY t.name, c.column_id", - table_name - ) - } - CatalogProc::SprocColumns => { - let proc_name = cat_req - .params.first() - .map(|p| p.value_sql.trim_matches('\'')) - .unwrap_or("%"); - format!( - "SELECT s.name AS PROCEDURE_OWNER, r.name AS PROCEDURE_NAME, p.name AS COLUMN_NAME, p.parameter_id AS ORDINAL_POSITION, ty.name AS TYPE_NAME FROM sys.parameters p JOIN sys.routines r ON p.object_id = r.object_id JOIN sys.schemas s ON r.schema_id = s.id JOIN sys.types ty ON p.user_type_id = ty.user_type_id WHERE r.name LIKE '{}' ORDER BY r.name, p.parameter_id", - proc_name - ) - } - CatalogProc::PrimaryKeys => { - let table_name = cat_req - .params.first() - .map(|p| p.value_sql.trim_matches('\'')) - .unwrap_or("%"); - format!( - "SELECT s.name AS TABLE_OWNER, t.name AS TABLE_NAME, c.name AS COLUMN_NAME, c.column_id AS KEY_SEQ, pk.name AS PK_NAME FROM sys.columns c JOIN sys.tables t ON c.object_id = t.object_id JOIN sys.schemas s ON t.schema_id = s.id JOIN (SELECT ic.object_id, ic.column_id, i.name FROM sys.index_columns ic JOIN sys.indexes i ON ic.object_id = i.object_id AND ic.index_id = i.index_id WHERE i.is_primary_key = 1) pk ON c.object_id = pk.object_id AND c.column_id = pk.column_id WHERE t.name LIKE '{}' ORDER BY t.name, c.column_id", - table_name - ) - } - CatalogProc::DescribeCursor => { - let cursor_handle = cat_req - .params - .get(2) - .and_then(|p| p.value_sql.parse::().ok()) - .unwrap_or(0); - if let Some(session_id) = self.session_id { - match self.db.cursor_rpc_fetch( - session_id, - cursor_handle, - 2, - 0, - 1, - ) { - Ok(fetch_result) => { - let mut buf = PacketBuilder::new(); - let col_names = vec![ - "reference_name".to_string(), - "cursor_name".to_string(), - "cursor_scope".to_string(), - "status".to_string(), - "model".to_string(), - "concurrency".to_string(), - "scrollable".to_string(), - "open_status".to_string(), - "cursor_rows".to_string(), - "fetch_status".to_string(), - "column_count".to_string(), - "row_count".to_string(), - "last_operation".to_string(), - "cursor_handle".to_string(), - ]; - let col_types: Vec<_> = col_names.iter() - .map(|_| super::tds::type_mapping::runtime_type_to_tds(&iridium_core::types::DataType::Int)) - .collect(); - let cursor_name = format!( - "#rpc_cursor_{}", - cursor_handle - ); - let row = vec![ - iridium_core::types::Value::NVarChar(cursor_name.clone()), - iridium_core::types::Value::NVarChar(cursor_name), - iridium_core::types::Value::Int(1), - iridium_core::types::Value::Int(0), - iridium_core::types::Value::Int(0), - iridium_core::types::Value::Int(1), - iridium_core::types::Value::Int(1), - iridium_core::types::Value::Int(if fetch_result.rows.is_empty() { 0 } else { 1 }), - iridium_core::types::Value::Int(fetch_result.rows.len() as i32), - iridium_core::types::Value::Int(fetch_result.fetch_status), - iridium_core::types::Value::Int(fetch_result.columns.len() as i32), - iridium_core::types::Value::Int(fetch_result.rows.len() as i32), - iridium_core::types::Value::Int(0), - iridium_core::types::Value::Int(cursor_handle), - ]; - tokens::write_colmetadata( - &mut buf, &col_names, &col_types, - ); - tokens::write_row( - &mut buf, &row, &col_types, 0, - ); - tokens::write_done( - &mut buf, - tokens::DONE_FINAL - | tokens::DONE_COUNT, - 0, - 1, - ); - let _ = packet::write_packet( - &mut writer, - TABULAR_RESULT, - buf.as_bytes(), - ) - .await; - } - Err(e) => { - let err_resp = build_error_response(&e); - let _ = packet::write_packet( - &mut writer, - TABULAR_RESULT, - &err_resp.data, - ) - .await; - } - } - } else { - let err = DbError::Execution( - "no session for cursor operation".into(), - ); - let err_resp = build_error_response(&err); - let _ = packet::write_packet( - &mut writer, - TABULAR_RESULT, - &err_resp.data, - ) - .await; - } - String::new() - } - }; - if !sql.is_empty() { - match self.execute_sql(&sql, &mut writer).await { - Ok(_) => {} - Err(e) => { - log::error!("Catalog RPC error: {}", e); - let err_resp = build_error_response(&e); - let _ = packet::write_packet( - &mut writer, - TABULAR_RESULT, - &err_resp.data, - ) - .await; - } - } - } - } - } - } - Ok(None) => { - let preview_len = data.len().min(96); - log::warn!( - "[conn={}] Unsupported RPC request ({} bytes), first {} bytes: {:02X?}", - self.connection_id, - data.len(), - preview_len, - &data[..preview_len] - ); - let err = DbError::Parse("unsupported RPC request".into()); - let err_resp = build_error_response(&err); - let _ = packet::write_packet( - &mut writer, - TABULAR_RESULT, - &err_resp.data, - ) - .await; - } - Err(e) => { - let err = iridium_core::error::DbError::Parse(e.to_string()); - let err_resp = build_error_response(&err); - let _ = packet::write_packet( - &mut writer, - TABULAR_RESULT, - &err_resp.data, - ) - .await; - } - }, - BULK_LOAD => { - log::info!("[conn={}] BULK_LOAD received", self.connection_id); - if let Some(session_id) = self.session_id { - let (active, target, columns, received_metadata) = self - .db - .session_options(session_id) - .map(|_opts| self.db.get_bulk_load_state(session_id)) - .unwrap_or((false, None, None, false)); - - if active && target.is_some() && columns.is_some() { - let target = target.unwrap(); - let columns = columns.unwrap(); - - let mut reader = packet::PacketReader::new(&data); - let mut column_types = Vec::new(); - - if !received_metadata { - let token = reader.read_u8().map_err(|e| e.to_string())?; - if token != tokens::COLMETADATA_TOKEN { - return Err(format!( - "Expected COLMETADATA (0x81), got 0x{:02X}", - token - )); - } - let count = - reader.read_u16_le().map_err(|e| e.to_string())? - as usize; - for _ in 0..count { - reader.skip(4).map_err(|e| e.to_string())?; // UserType - let _flags = - reader.read_u16_le().map_err(|e| e.to_string())?; - let ti = super::tds::type_mapping::read_type_info( - &mut reader, - ) - .map_err(|e| e.to_string())?; - column_types.push(ti); - let name_len = - reader.read_u8().map_err(|e| e.to_string())? - as usize; - let _name = reader - .read_utf16le(name_len) - .map_err(|e| e.to_string())?; - } - // Mark metadata as received - self.db - .set_bulk_load_active( - session_id, - true, - target.clone(), - columns.clone(), - true, - ) - .map_err(|e| e.to_string())?; - } - - // For now, we still assume the first packet has both metadata and some rows, - // or we'd need to store column_types in the session state to handle subsequent row-only packets. - - match parse_bulk_load_data(&data, &columns) { - Ok(bulk_data) => { - log::info!( - "[conn={}] Parsed {} bulk rows for table {}", - self.connection_id, - bulk_data.rows.len(), - target.name - ); - - // Construct INSERT statements or use MutationExecutor directly. - // For simplicity, we'll build a batch of INSERTs. - let mut sql = String::new(); - let col_names: Vec = - columns.iter().map(|c| c.name.clone()).collect(); - let col_list = col_names.join(", "); - - for row in bulk_data.rows { - let vals: Vec = row - .iter() - .map(|v| v.to_sql_literal()) - .collect(); - sql.push_str(&format!( - "INSERT INTO {}.{} ({}) VALUES ({});\n", - target.schema_or_dbo(), - target.name, - col_list, - vals.join(", ") - )); - } - - // Reset bulk state before executing to avoid recursion/loops - self.db - .set_bulk_load_active( - session_id, - false, - target.clone(), - columns.clone(), - false, - ) - .map_err(|e| e.to_string())?; - - match self.execute_sql(&sql, &mut writer).await { - Ok(_) => {} - Err(e) => { - log::error!( - "Bulk insert execution error: {}", - e - ); - let err_resp = build_error_response(&e); - let _ = packet::write_packet( - &mut writer, - TABULAR_RESULT, - &err_resp.data, - ) - .await; - } - } - } - Err(e) => { - log::error!("Bulk data parse error: {}", e); - let err = DbError::Parse(e.to_string()); - let err_resp = build_error_response(&err); - let _ = packet::write_packet( - &mut writer, - TABULAR_RESULT, - &err_resp.data, - ) - .await; - } - } - } else { - log::warn!("[conn={}] Received BULK_LOAD but bulk load was not expected", self.connection_id); - } - } - } - ATTENTION => { - log::debug!("[conn={}] ATTENTION received", self.connection_id); - let mut attn = PacketBuilder::new(); - tokens::write_done(&mut attn, tokens::DONE_ATTN, 0, 0); - let _ = - packet::write_packet(&mut writer, TABULAR_RESULT, attn.as_bytes()) - .await; - } - _ => { - log::warn!( - "[conn={}] Unsupported packet type 0x{:02X}", - self.connection_id, - header.packet_type - ); - } - } - } - Err(ref e) if e.kind() == std::io::ErrorKind::UnexpectedEof => { - log::debug!("[conn={}] Client disconnected", self.connection_id); - break; - } - Err(e) => { - log::error!("[conn={}] Read error: {}", self.connection_id, e); - break; - } - } - } - - if let Some(sid) = self.session_id.take() { - self.session_pool.checkin(self.db.as_ref(), sid); - } - - Ok(()) - } - - async fn handle_sql_batch( - &mut self, - data: &[u8], - writer: &mut W, - ) -> Result { - let sql = match parse_sql_batch(data) { - Ok(s) => s, - Err(e) => { - let err = iridium_core::error::DbError::Parse(e.to_string()); - let err_resp = build_error_response(&err); - packet::write_packet(writer, TABULAR_RESULT, &err_resp.data) - .await - .map_err(|e| iridium_core::error::DbError::Execution(e.to_string()))?; - return Ok(true); - } - }; - - if !sql.trim().is_empty() { - log::info!( - "[conn={}] SQL batch received:\n{}", - self.connection_id, - format_sql_for_log(sql.trim()) - ); - } - self.execute_sql(sql.trim(), writer).await - } - - async fn execute_sql( - &mut self, - sql: &str, - writer: &mut W, - ) -> Result { - if sql.is_empty() { - let mut b = PacketBuilder::new(); - tokens::write_done(&mut b, tokens::DONE_FINAL, 1, 0); - packet::write_packet(writer, TABULAR_RESULT, b.as_bytes()) - .await - .map_err(|e| iridium_core::error::DbError::Execution(e.to_string()))?; - return Ok(true); - } - - let session_id = self.session_id.ok_or_else(|| { - iridium_core::error::DbError::Execution("session not initialized".to_string()) - })?; - - if is_ssms_contained_auth_probe(sql) { - if let Some(db_name) = extract_leading_use_database(sql) { - self.apply_use_database(session_id, &db_name, writer) - .await?; - } - // SSMS expects a scalar response for this probe; returning 0 keeps the flow compatible. - let data = build_single_int_result("", 0); - packet::write_packet(writer, TABULAR_RESULT, &data) - .await - .map_err(|e| iridium_core::error::DbError::Execution(e.to_string()))?; - return Ok(true); - } - - if let Some(db_name) = parse_simple_use_database(sql) { - self.apply_use_database(session_id, &db_name, writer) - .await?; - return Ok(true); - } - - log_sql_execution(self.connection_id, sql); - let force_sysdac_probe_int = self::compat::is_sysdac_instances_probe(sql); - match self.db.execute_session_batch_sql_multi(session_id, sql) { - Ok(results) => { - let count = results.len(); - let mut b = PacketBuilder::with_capacity(4096); - let textsize = self - .db - .session_options(session_id) - .map(|opts| opts.textsize.max(0) as usize) - .unwrap_or(4096); - - for (i, result) in results.into_iter().enumerate() { - let is_last = i == count - 1; - - match result { - Some(mut query_result) => { - let is_proc = query_result.is_procedure; - let return_status = query_result.return_status; - - if !query_result.columns.is_empty() { - if force_sysdac_probe_int - && query_result.columns.len() == 1 - && query_result.rows.len() == 1 - { - query_result.column_types[0] = DataType::Int; - if let Some(row) = query_result.rows.get_mut(0) { - if let Some(value) = row.get_mut(0) { - let int_val = match &*value { - Value::Null => 0, - other => other.to_integer_i64().unwrap_or(0) as i32, - }; - *value = Value::Int(int_val); - } - } - } - - let mut types = Vec::new(); - log::debug!( - "[conn={}] Result set: columns={}, types={}", - self.connection_id, - query_result.columns.len(), - query_result.column_types.len() - ); - for ct in &query_result.column_types { - types.push(crate::tds::type_mapping::runtime_type_to_tds(ct)); - } - for (idx, col_name) in query_result.columns.iter().enumerate() { - if let (Some(runtime_ty), Some(tds_ty)) = - (query_result.column_types.get(idx), types.get(idx)) - { - log::debug!( - "[conn={}] COLMETADATA[{}]: name='{}' runtime={:?} tds=0x{:02X} len={:02X?}", - self.connection_id, - idx, - col_name, - runtime_ty, - tds_ty.tds_type, - tds_ty.length_prefix - ); - } - } - tokens::write_colmetadata(&mut b, &query_result.columns, &types); - for row in &query_result.rows { - tokens::write_row(&mut b, row, &types, textsize); - } - - if is_proc { - tokens::write_done_in_proc( - &mut b, - tokens::DONE_MORE | tokens::DONE_COUNT, - 1, - query_result.rows.len() as u64, - ); - } else { - let done_status = if is_last && return_status.is_none() { - tokens::DONE_FINAL - } else { - tokens::DONE_MORE - }; - tokens::write_done( - &mut b, - done_status | tokens::DONE_COUNT, - 1, - query_result.rows.len() as u64, - ); - } - } else if !is_proc { - let done_status = if is_last && return_status.is_none() { - tokens::DONE_FINAL - } else { - tokens::DONE_MORE - }; - tokens::write_done( - &mut b, - done_status | tokens::DONE_COUNT, - 1, - query_result.rows.len() as u64, - ); - } - - if let Some(code) = return_status { - tokens::write_returnstatus(&mut b, code); - let done_status = if is_last { - tokens::DONE_FINAL - } else { - tokens::DONE_MORE - }; - tokens::write_doneproc(&mut b, done_status, 1, 0); - } - } - None => { - let done_status = if is_last { - tokens::DONE_FINAL - } else { - tokens::DONE_MORE - }; - tokens::write_done(&mut b, done_status, 1, 0); - } - } - } - - // If no statements at all, send a final DONE - if count == 0 { - tokens::write_done(&mut b, tokens::DONE_FINAL, 1, 0); - } - - packet::write_packet(writer, TABULAR_RESULT, b.as_bytes()) - .await - .map_err(|e| iridium_core::error::DbError::Execution(e.to_string()))?; - } - Err(e) => { - log::warn!( - "[conn={}] SQL execution failed for batch:\n{}\nerror: {}", - self.connection_id, - format_sql_for_log(sql), - e - ); - let err_resp = build_error_response(&e); - packet::write_packet(writer, TABULAR_RESULT, &err_resp.data) - .await - .map_err(|e| iridium_core::error::DbError::Execution(e.to_string()))?; - } - } - - Ok(true) - } - - async fn apply_use_database( - &mut self, - session_id: SessionId, - db_name: &str, - writer: &mut W, - ) -> Result<(), iridium_core::error::DbError> { - let old_db = self.database.clone(); - self.database = db_name.to_string(); - if let Err(e) = self - .db - .set_session_database(session_id, self.database.clone()) - { - log::error!( - "[conn={}] Failed to update session database context: {}", - self.connection_id, - e - ); - } - - let data = build_use_database_response(&self.database, &old_db); - packet::write_packet(writer, TABULAR_RESULT, &data) - .await - .map_err(|e| iridium_core::error::DbError::Execution(e.to_string())) - } } static SQL_KEYWORD_RE: Lazy = Lazy::new(|| { @@ -1443,7 +66,7 @@ static SQL_KEYWORD_RE: Lazy = Lazy::new(|| { .expect("valid SQL keyword regex") }); -fn format_sql_for_log(sql: &str) -> String { +pub(crate) fn format_sql_for_log(sql: &str) -> String { let text = if std::env::var("IRIDIUM_LOG_FULL_SQL") .or_else(|_| std::env::var("TSQL_LOG_FULL_SQL")) .map(|v| v == "1" || v.eq_ignore_ascii_case("true")) @@ -1469,7 +92,7 @@ fn format_sql_for_log(sql: &str) -> String { .to_string() } -fn log_sql_execution(connection_id: u64, sql: &str) { +pub(crate) fn log_sql_execution(connection_id: u64, sql: &str) { if false { return; } @@ -1481,7 +104,7 @@ fn log_sql_execution(connection_id: u64, sql: &str) { ); } -fn truncate_for_log(text: &str, max_chars: usize) -> String { +pub(crate) fn truncate_for_log(text: &str, max_chars: usize) -> String { let total_chars = text.chars().count(); if total_chars <= max_chars { return text.to_string(); @@ -1497,7 +120,7 @@ fn truncate_for_log(text: &str, max_chars: usize) -> String { preview } -fn log_packet(connection_id: u64, stage: &str, header: &packet::PacketHeader, data: &[u8]) { +pub(crate) fn log_packet(connection_id: u64, stage: &str, header: &packet::PacketHeader, data: &[u8]) { let preview_len = data.len().min(96); log::debug!( "[conn={}] {} packet type=0x{:02X} len={} preview={:02X?}", diff --git a/crates/iridium_server/src/tds/rpc/parser/mod.rs b/crates/iridium_server/src/tds/rpc/parser/mod.rs new file mode 100644 index 0000000..98c47e3 --- /dev/null +++ b/crates/iridium_server/src/tds/rpc/parser/mod.rs @@ -0,0 +1,110 @@ +pub mod types; +pub mod utils; +pub mod parser; + +pub use types::*; +pub use utils::*; +pub use parser::*; + +#[cfg(test)] +mod tests { + use super::*; + + #[test] + fn rpc_proc_resolves_by_id() { + // MS-TDS spec IDs (corrected) + assert_eq!(RpcProc::from_id(1), Some(RpcProc::CursorOpen)); + assert_eq!(RpcProc::from_id(2), Some(RpcProc::CursorOpen)); + assert_eq!(RpcProc::from_id(3), Some(RpcProc::CursorPrepare)); // sp_cursorprepare + assert_eq!(RpcProc::from_id(4), Some(RpcProc::CursorExecute)); // sp_cursorexecute + assert_eq!(RpcProc::from_id(5), Some(RpcProc::CursorPrepExec)); // sp_cursorprepexec + assert_eq!(RpcProc::from_id(6), Some(RpcProc::CursorUnprepare)); // sp_cursorunprepare + assert_eq!(RpcProc::from_id(7), Some(RpcProc::CursorFetch)); + assert_eq!(RpcProc::from_id(8), Some(RpcProc::CursorOption)); // sp_cursoroption + assert_eq!(RpcProc::from_id(9), Some(RpcProc::CursorClose)); // sp_cursorclose + assert_eq!(RpcProc::from_id(10), Some(RpcProc::ExecuteSql)); + assert_eq!(RpcProc::from_id(11), Some(RpcProc::Prepare)); // sp_prepare + assert_eq!(RpcProc::from_id(12), Some(RpcProc::Execute)); // sp_execute + assert_eq!(RpcProc::from_id(13), Some(RpcProc::PrepExec)); // sp_prepexec + assert_eq!(RpcProc::from_id(14), Some(RpcProc::PrepExecRpc)); // sp_prepexecrpc + assert_eq!(RpcProc::from_id(15), Some(RpcProc::Unprepare)); // sp_unprepare + assert_eq!(RpcProc::from_id(42), None); + } + + #[test] + fn rpc_proc_resolves_by_name() { + assert_eq!(RpcProc::from_name("sp_cursor"), Some(RpcProc::CursorOpen)); + assert_eq!( + RpcProc::from_name("sp_cursoropen"), + Some(RpcProc::CursorOpen) + ); + assert_eq!( + RpcProc::from_name("sp_cursorclose"), + Some(RpcProc::CursorClose) + ); + assert_eq!( + RpcProc::from_name("sp_cursorfetch"), + Some(RpcProc::CursorFetch) + ); + assert_eq!( + RpcProc::from_name("[dbo].[sp_cursorprepare]"), + Some(RpcProc::CursorPrepare) + ); + assert_eq!( + RpcProc::from_name("sp_cursorexecute"), + Some(RpcProc::CursorExecute) + ); + assert_eq!( + RpcProc::from_name("sp_cursorunprepare"), + Some(RpcProc::CursorUnprepare) + ); + assert_eq!( + RpcProc::from_name("sp_cursoroption"), + Some(RpcProc::CursorOption) + ); + assert_eq!( + RpcProc::from_name("sp_executesql"), + Some(RpcProc::ExecuteSql) + ); + assert_eq!(RpcProc::from_name("sp_prepexec"), Some(RpcProc::PrepExec)); + assert_eq!(RpcProc::from_name("sp_prepare"), Some(RpcProc::Prepare)); + assert_eq!(RpcProc::from_name("sp_execute"), Some(RpcProc::Execute)); + assert_eq!(RpcProc::from_name("sp_unprepare"), Some(RpcProc::Unprepare)); + assert_eq!( + RpcProc::from_name("sp_reset_connection"), + Some(RpcProc::ResetConnection) + ); + assert_eq!(RpcProc::from_name("sp_help"), None); + } + + #[test] + fn rpc_proc_is_cursor() { + assert!(RpcProc::CursorOpen.is_cursor()); + assert!(RpcProc::CursorClose.is_cursor()); + assert!(RpcProc::CursorFetch.is_cursor()); + assert!(RpcProc::CursorPrepare.is_cursor()); + assert!(RpcProc::CursorExecute.is_cursor()); + assert!(RpcProc::CursorUnprepare.is_cursor()); + assert!(RpcProc::CursorOption.is_cursor()); + assert!(!RpcProc::ExecuteSql.is_cursor()); + assert!(!RpcProc::PrepExec.is_cursor()); + assert!(!RpcProc::Prepare.is_cursor()); + assert!(!RpcProc::Execute.is_cursor()); + assert!(!RpcProc::Unprepare.is_cursor()); + assert!(!RpcProc::ResetConnection.is_cursor()); + } + + #[test] + fn param_decl_parsing() { + let decls = parse_param_decl("@x INT, @name NVARCHAR(50)"); + assert_eq!(decls.len(), 2); + assert_eq!(decls[0], ("@x".to_string(), "INT".to_string())); + assert_eq!(decls[1], ("@name".to_string(), "NVARCHAR(50)".to_string())); + + let empty = parse_param_decl(""); + assert!(empty.is_empty()); + + let single = parse_param_decl("@y INT"); + assert_eq!(single.len(), 1); + } +} diff --git a/crates/iridium_server/src/tds/rpc/parser.rs b/crates/iridium_server/src/tds/rpc/parser/parser.rs similarity index 64% rename from crates/iridium_server/src/tds/rpc/parser.rs rename to crates/iridium_server/src/tds/rpc/parser/parser.rs index 50975f0..843ec32 100644 --- a/crates/iridium_server/src/tds/rpc/parser.rs +++ b/crates/iridium_server/src/tds/rpc/parser/parser.rs @@ -1,231 +1,26 @@ use crate::tds::packet::PacketReader; use std::io; +use super::super::decode; +use super::super::types::RpcParam; +use super::types::*; +use super::utils::*; -use super::decode; -use super::types::RpcParam; - -#[derive(Debug, Clone)] -enum RpcProcSelector { - Id(u16), - Name(String), -} - -struct RpcFrameParser<'a> { +pub struct RpcFrameParser<'a> { reader: PacketReader<'a>, } -#[derive(Debug, Clone)] -pub enum RpcRequest { - Sql(SqlRpcRequest), - Cursor(CursorRpcRequest), - Prepare(PrepareRpcRequest), - Execute(ExecuteRpcRequest), - Unprepare(UnprepareRpcRequest), - PrepExec(PrepExecRpcRequest), - ResetConnection, - Catalog(CatalogRpcRequest), -} - -#[derive(Debug, Clone, Copy, PartialEq, Eq)] -pub enum CatalogProc { - Tables, - Columns, - SprocColumns, - PrimaryKeys, - DescribeCursor, -} - -#[derive(Debug, Clone)] -pub struct CatalogRpcRequest { - pub proc: CatalogProc, - pub params: Vec, -} - -#[derive(Debug, Clone)] -pub struct SqlRpcRequest { - pub sql: String, - pub params: Vec, -} - -#[derive(Debug, Clone)] -pub struct PrepareRpcRequest { - pub stmt_handle: Option, - pub sql: String, - pub param_decl: String, - pub params: Vec, -} - -#[derive(Debug, Clone)] -pub struct ExecuteRpcRequest { - pub stmt_handle: i32, - pub params: Vec, -} - -#[derive(Debug, Clone)] -pub struct UnprepareRpcRequest { - pub stmt_handle: i32, -} - -#[derive(Debug, Clone)] -pub struct PrepExecRpcRequest { - pub stmt_handle: Option, - pub sql: String, - pub param_decl: String, - pub params: Vec, -} - -#[derive(Debug, Clone)] -pub struct CursorRpcRequest { - pub cursor_op: CursorOp, - pub cursor_handle: Option, - pub scroll_opt: Option, - pub cc_opt: Option, - pub row_count: Option, - pub sql: Option, - pub param_def: Option, - pub params: Vec, - pub fetch_type: Option, - pub row_num: Option, - pub n_rows: Option, -} - -#[derive(Debug, Clone, Copy, PartialEq, Eq)] -pub enum CursorOp { - Open, - Fetch, - Close, - Prepare, - Execute, - PrepExec, - Unprepare, - Option, -} - -#[derive(Debug, Clone, Copy, PartialEq, Eq)] -pub enum RpcProc { - CursorOpen, - CursorClose, - CursorFetch, - CursorPrepare, - CursorExecute, - CursorPrepExec, - CursorUnprepare, - CursorOption, - ExecuteSql, - PrepExec, - PrepExecRpc, - Prepare, - Execute, - Unprepare, - ResetConnection, - SpTables, - SpColumns, - SpSprocColumns, - SpPkeys, - SpDescribeCursor, -} - -impl RpcProc { - fn from_id(id: u16) -> Option { - match id { - 1 => Some(RpcProc::CursorOpen), // sp_cursor - 2 => Some(RpcProc::CursorOpen), // sp_cursoropen - 3 => Some(RpcProc::CursorPrepare), // sp_cursorprepare (MS-TDS) - 4 => Some(RpcProc::CursorExecute), // sp_cursorexecute (MS-TDS) - 5 => Some(RpcProc::CursorPrepExec), // sp_cursorprepexec (MS-TDS) - 6 => Some(RpcProc::CursorUnprepare), // sp_cursorunprepare (MS-TDS) - 7 => Some(RpcProc::CursorFetch), // sp_cursorfetch - 8 => Some(RpcProc::CursorOption), // sp_cursoroption (MS-TDS) - 9 => Some(RpcProc::CursorClose), // sp_cursorclose (MS-TDS) - 10 => Some(RpcProc::ExecuteSql), // sp_executesql - 11 => Some(RpcProc::Prepare), // sp_prepare (MS-TDS) - 12 => Some(RpcProc::Execute), // sp_execute (MS-TDS) - 13 => Some(RpcProc::PrepExec), // sp_prepexec - 14 => Some(RpcProc::PrepExecRpc), // sp_prepexecrpc (MS-TDS) - 15 => Some(RpcProc::Unprepare), // sp_unprepare (MS-TDS) - _ => None, - } - } - - fn from_name(name: &str) -> Option { - let n = normalize_proc_name(name); - match n.as_str() { - "sp_cursor" | "sp_cursoropen" => Some(RpcProc::CursorOpen), - "sp_cursorclose" => Some(RpcProc::CursorClose), - "sp_cursorfetch" => Some(RpcProc::CursorFetch), - "sp_cursorprepare" => Some(RpcProc::CursorPrepare), - "sp_cursorexecute" => Some(RpcProc::CursorExecute), - "sp_cursorprepexec" => Some(RpcProc::CursorPrepExec), - "sp_cursorunprepare" => Some(RpcProc::CursorUnprepare), - "sp_cursoroption" => Some(RpcProc::CursorOption), - "sp_executesql" => Some(RpcProc::ExecuteSql), - "sp_prepexec" => Some(RpcProc::PrepExec), - "sp_prepexecrpc" => Some(RpcProc::PrepExecRpc), - "sp_prepare" => Some(RpcProc::Prepare), - "sp_execute" => Some(RpcProc::Execute), - "sp_unprepare" => Some(RpcProc::Unprepare), - "sp_reset_connection" => Some(RpcProc::ResetConnection), - "sp_tables" => Some(RpcProc::SpTables), - "sp_columns" => Some(RpcProc::SpColumns), - "sp_sproc_columns" => Some(RpcProc::SpSprocColumns), - "sp_pkeys" => Some(RpcProc::SpPkeys), - "sp_describe_cursor" => Some(RpcProc::SpDescribeCursor), - _ => None, - } - } - - pub fn is_cursor(&self) -> bool { - matches!( - self, - RpcProc::CursorOpen - | RpcProc::CursorClose - | RpcProc::CursorFetch - | RpcProc::CursorPrepare - | RpcProc::CursorExecute - | RpcProc::CursorPrepExec - | RpcProc::CursorUnprepare - | RpcProc::CursorOption - ) - } -} - -impl CursorOp { - #[allow(dead_code)] - fn from_id(_id: u16) -> Option { - None - } - - #[allow(dead_code)] - fn from_name(name: &str) -> Option { - let n = normalize_proc_name(name); - match n.as_str() { - "sp_cursor" | "sp_cursoropen" => Some(CursorOp::Open), - "sp_cursorclose" => Some(CursorOp::Close), - "sp_cursorfetch" => Some(CursorOp::Fetch), - "sp_cursorprepare" => Some(CursorOp::Prepare), - "sp_cursorexecute" => Some(CursorOp::Execute), - "sp_cursorprepexec" => Some(CursorOp::PrepExec), - "sp_cursorunprepare" => Some(CursorOp::Unprepare), - "sp_cursoroption" => Some(CursorOp::Option), - _ => None, - } - } -} - -/// Parse an RPC request packet (sp_executesql=10, sp_prepexec=13, cursor procedures). -/// Returns None if not a supported RPC call. pub fn parse_rpc(data: &[u8]) -> io::Result> { RpcFrameParser::new(data).parse() } impl<'a> RpcFrameParser<'a> { - fn new(data: &'a [u8]) -> Self { + pub(crate) fn new(data: &'a [u8]) -> Self { Self { reader: PacketReader::new(data), } } - fn parse(mut self) -> io::Result> { + pub(crate) fn parse(mut self) -> io::Result> { self.skip_all_headers()?; let proc_selector = self.read_proc_selector()?; self.skip_rpc_flags()?; @@ -363,7 +158,7 @@ impl<'a> RpcFrameParser<'a> { fn parse_execute_rpc(&mut self) -> io::Result> { let mut stmt_handle: Option = None; - let params: Vec = vec![]; + let mut params: Vec = vec![]; let mut param_idx: usize = 0; while self.reader.remaining() > 0 { @@ -382,7 +177,16 @@ impl<'a> RpcFrameParser<'a> { match param_idx { 0 => stmt_handle = self.read_int_param(type_id)?, - _ => self.skip_typed_value(type_id)?, + _ => { + // Collect parameters if any + let decoded = decode::read_typed_value(&mut self.reader, type_id)?; + params.push(RpcParam { + name: String::new(), + type_name: decoded.type_name, + value_sql: decoded.value_sql, + tvp_rows: decoded.tvp_rows, + }); + } } param_idx += 1; } @@ -537,7 +341,15 @@ impl<'a> RpcFrameParser<'a> { CursorOp::Execute => match param_idx { 0 => cursor_handle = self.read_int_param(type_id)?, 1 => row_count = self.read_int_param(type_id)?, - _ => self.skip_typed_value(type_id)?, + _ => { + let decoded = decode::read_typed_value(&mut self.reader, type_id)?; + params.push(RpcParam { + name: String::new(), + type_name: decoded.type_name, + value_sql: decoded.value_sql, + tvp_rows: decoded.tvp_rows, + }); + } }, CursorOp::PrepExec => match param_idx { 0 => {} @@ -581,7 +393,6 @@ impl<'a> RpcFrameParser<'a> { fn read_int_param(&mut self, type_id: u8) -> io::Result> { match type_id { 0x26 => { - // INTN (4 bytes) - reads as i32 from little-endian u32 if self.reader.remaining() < 4 { return Ok(None); } @@ -594,7 +405,6 @@ impl<'a> RpcFrameParser<'a> { } } 0x38 => { - // BIGINT - reads as i32 from little-endian u64 if self.reader.remaining() < 8 { return Ok(None); } @@ -610,33 +420,26 @@ impl<'a> RpcFrameParser<'a> { fn skip_typed_value(&mut self, type_id: u8) -> io::Result<()> { match type_id { - 0x1F => { - // NULL - Ok(()) - } + 0x1F => Ok(()), 0x26 => { - // INTN if self.reader.remaining() >= 5 { self.reader.skip(5)?; } Ok(()) } 0x38 => { - // BIGINTN if self.reader.remaining() >= 9 { self.reader.skip(9)?; } Ok(()) } 0x6A | 0x68 => { - // INT | SMALLINT if self.reader.remaining() >= 1 { self.reader.skip(1)?; } Ok(()) } _ => { - // Skip based on type family let skip_chars = match type_id { 0xE7 | 0xA7 => 9, 0x22 | 0x21 => 5, @@ -788,7 +591,6 @@ impl<'a> RpcFrameParser<'a> { } } 0x63 | 0x23 => { - // NTEXT or TEXT if self.reader.remaining() < 13 { return Ok(None); } @@ -830,148 +632,3 @@ impl<'a> RpcFrameParser<'a> { } } } - -#[allow(dead_code)] -fn is_supported_rpc_proc(proc: &RpcProcSelector) -> bool { - match proc { - RpcProcSelector::Id(id) => *id == 10 || *id == 13, - RpcProcSelector::Name(name) => { - let base = normalize_proc_name(name); - base == "sp_executesql" - || base == "sp_prepexec" - || base == "sp_prepare" - || base == "sp_unprepare" - } - } -} - -fn normalize_proc_name(name: &str) -> String { - let mut part = name.trim(); - if let Some(last) = part.rsplit('.').next() { - part = last; - } - part.trim_matches(|c| c == '[' || c == ']' || c == ' ') - .to_ascii_lowercase() -} - -pub fn parse_param_decl(decl: &str) -> Vec<(String, String)> { - if decl.trim().is_empty() { - return vec![]; - } - decl.split(',') - .filter_map(|part| { - let part = part.trim(); - let mut iter = part.splitn(2, char::is_whitespace); - let name = iter.next()?.trim().to_string(); - let type_name = iter.next()?.trim().to_string(); - if name.is_empty() || type_name.is_empty() { - None - } else { - Some((name, type_name)) - } - }) - .collect() -} - -#[cfg(test)] -mod tests { - use super::*; - - #[test] - fn rpc_proc_resolves_by_id() { - // MS-TDS spec IDs (corrected) - assert_eq!(RpcProc::from_id(1), Some(RpcProc::CursorOpen)); - assert_eq!(RpcProc::from_id(2), Some(RpcProc::CursorOpen)); - assert_eq!(RpcProc::from_id(3), Some(RpcProc::CursorPrepare)); // sp_cursorprepare - assert_eq!(RpcProc::from_id(4), Some(RpcProc::CursorExecute)); // sp_cursorexecute - assert_eq!(RpcProc::from_id(5), Some(RpcProc::CursorPrepExec)); // sp_cursorprepexec - assert_eq!(RpcProc::from_id(6), Some(RpcProc::CursorUnprepare)); // sp_cursorunprepare - assert_eq!(RpcProc::from_id(7), Some(RpcProc::CursorFetch)); - assert_eq!(RpcProc::from_id(8), Some(RpcProc::CursorOption)); // sp_cursoroption - assert_eq!(RpcProc::from_id(9), Some(RpcProc::CursorClose)); // sp_cursorclose - assert_eq!(RpcProc::from_id(10), Some(RpcProc::ExecuteSql)); - assert_eq!(RpcProc::from_id(11), Some(RpcProc::Prepare)); // sp_prepare - assert_eq!(RpcProc::from_id(12), Some(RpcProc::Execute)); // sp_execute - assert_eq!(RpcProc::from_id(13), Some(RpcProc::PrepExec)); // sp_prepexec - assert_eq!(RpcProc::from_id(14), Some(RpcProc::PrepExecRpc)); // sp_prepexecrpc - assert_eq!(RpcProc::from_id(15), Some(RpcProc::Unprepare)); // sp_unprepare - assert_eq!(RpcProc::from_id(42), None); - } - - #[test] - fn rpc_proc_resolves_by_name() { - assert_eq!(RpcProc::from_name("sp_cursor"), Some(RpcProc::CursorOpen)); - assert_eq!( - RpcProc::from_name("sp_cursoropen"), - Some(RpcProc::CursorOpen) - ); - assert_eq!( - RpcProc::from_name("sp_cursorclose"), - Some(RpcProc::CursorClose) - ); - assert_eq!( - RpcProc::from_name("sp_cursorfetch"), - Some(RpcProc::CursorFetch) - ); - assert_eq!( - RpcProc::from_name("[dbo].[sp_cursorprepare]"), - Some(RpcProc::CursorPrepare) - ); - assert_eq!( - RpcProc::from_name("sp_cursorexecute"), - Some(RpcProc::CursorExecute) - ); - assert_eq!( - RpcProc::from_name("sp_cursorunprepare"), - Some(RpcProc::CursorUnprepare) - ); - assert_eq!( - RpcProc::from_name("sp_cursoroption"), - Some(RpcProc::CursorOption) - ); - assert_eq!( - RpcProc::from_name("sp_executesql"), - Some(RpcProc::ExecuteSql) - ); - assert_eq!(RpcProc::from_name("sp_prepexec"), Some(RpcProc::PrepExec)); - assert_eq!(RpcProc::from_name("sp_prepare"), Some(RpcProc::Prepare)); - assert_eq!(RpcProc::from_name("sp_execute"), Some(RpcProc::Execute)); - assert_eq!(RpcProc::from_name("sp_unprepare"), Some(RpcProc::Unprepare)); - assert_eq!( - RpcProc::from_name("sp_reset_connection"), - Some(RpcProc::ResetConnection) - ); - assert_eq!(RpcProc::from_name("sp_help"), None); - } - - #[test] - fn rpc_proc_is_cursor() { - assert!(RpcProc::CursorOpen.is_cursor()); - assert!(RpcProc::CursorClose.is_cursor()); - assert!(RpcProc::CursorFetch.is_cursor()); - assert!(RpcProc::CursorPrepare.is_cursor()); - assert!(RpcProc::CursorExecute.is_cursor()); - assert!(RpcProc::CursorUnprepare.is_cursor()); - assert!(RpcProc::CursorOption.is_cursor()); - assert!(!RpcProc::ExecuteSql.is_cursor()); - assert!(!RpcProc::PrepExec.is_cursor()); - assert!(!RpcProc::Prepare.is_cursor()); - assert!(!RpcProc::Execute.is_cursor()); - assert!(!RpcProc::Unprepare.is_cursor()); - assert!(!RpcProc::ResetConnection.is_cursor()); - } - - #[test] - fn param_decl_parsing() { - let decls = parse_param_decl("@x INT, @name NVARCHAR(50)"); - assert_eq!(decls.len(), 2); - assert_eq!(decls[0], ("@x".to_string(), "INT".to_string())); - assert_eq!(decls[1], ("@name".to_string(), "NVARCHAR(50)".to_string())); - - let empty = parse_param_decl(""); - assert!(empty.is_empty()); - - let single = parse_param_decl("@y INT"); - assert_eq!(single.len(), 1); - } -} diff --git a/crates/iridium_server/src/tds/rpc/parser/types.rs b/crates/iridium_server/src/tds/rpc/parser/types.rs new file mode 100644 index 0000000..e8664b3 --- /dev/null +++ b/crates/iridium_server/src/tds/rpc/parser/types.rs @@ -0,0 +1,118 @@ +use super::super::types::RpcParam; + +#[derive(Debug, Clone)] +pub(crate) enum RpcProcSelector { + Id(u16), + Name(String), +} + +#[derive(Debug, Clone)] +pub enum RpcRequest { + Sql(SqlRpcRequest), + Cursor(CursorRpcRequest), + Prepare(PrepareRpcRequest), + Execute(ExecuteRpcRequest), + Unprepare(UnprepareRpcRequest), + PrepExec(PrepExecRpcRequest), + ResetConnection, + Catalog(CatalogRpcRequest), +} + +#[derive(Debug, Clone, Copy, PartialEq, Eq)] +pub enum CatalogProc { + Tables, + Columns, + SprocColumns, + PrimaryKeys, + DescribeCursor, +} + +#[derive(Debug, Clone)] +pub struct CatalogRpcRequest { + pub proc: CatalogProc, + pub params: Vec, +} + +#[derive(Debug, Clone)] +pub struct SqlRpcRequest { + pub sql: String, + pub params: Vec, +} + +#[derive(Debug, Clone)] +pub struct PrepareRpcRequest { + pub stmt_handle: Option, + pub sql: String, + pub param_decl: String, + pub params: Vec, +} + +#[derive(Debug, Clone)] +pub struct ExecuteRpcRequest { + pub stmt_handle: i32, + pub params: Vec, +} + +#[derive(Debug, Clone)] +pub struct UnprepareRpcRequest { + pub stmt_handle: i32, +} + +#[derive(Debug, Clone)] +pub struct PrepExecRpcRequest { + pub stmt_handle: Option, + pub sql: String, + pub param_decl: String, + pub params: Vec, +} + +#[derive(Debug, Clone)] +pub struct CursorRpcRequest { + pub cursor_op: CursorOp, + pub cursor_handle: Option, + pub scroll_opt: Option, + pub cc_opt: Option, + pub row_count: Option, + pub sql: Option, + pub param_def: Option, + pub params: Vec, + pub fetch_type: Option, + pub row_num: Option, + pub n_rows: Option, +} + +#[derive(Debug, Clone, Copy, PartialEq, Eq)] +pub enum CursorOp { + Open, + Fetch, + Close, + Prepare, + Execute, + PrepExec, + Unprepare, + Option, +} + +#[derive(Debug, Clone, Copy, PartialEq, Eq)] +pub enum RpcProc { + CursorOpen, + CursorClose, + CursorFetch, + CursorPrepare, + CursorExecute, + CursorPrepExec, + CursorUnprepare, + CursorOption, + ExecuteSql, + PrepExec, + PrepExecRpc, + Prepare, + Execute, + Unprepare, + ResetConnection, + SpTables, + SpColumns, + SpSprocColumns, + SpPkeys, + SpDescribeCursor, +} diff --git a/crates/iridium_server/src/tds/rpc/parser/utils.rs b/crates/iridium_server/src/tds/rpc/parser/utils.rs new file mode 100644 index 0000000..8a6e82b --- /dev/null +++ b/crates/iridium_server/src/tds/rpc/parser/utils.rs @@ -0,0 +1,130 @@ +use super::types::{RpcProc, CursorOp, RpcProcSelector}; + +impl RpcProc { + pub(crate) fn from_id(id: u16) -> Option { + match id { + 1 => Some(RpcProc::CursorOpen), // sp_cursor + 2 => Some(RpcProc::CursorOpen), // sp_cursoropen + 3 => Some(RpcProc::CursorPrepare), // sp_cursorprepare (MS-TDS) + 4 => Some(RpcProc::CursorExecute), // sp_cursorexecute (MS-TDS) + 5 => Some(RpcProc::CursorPrepExec), // sp_cursorprepexec (MS-TDS) + 6 => Some(RpcProc::CursorUnprepare), // sp_cursorunprepare (MS-TDS) + 7 => Some(RpcProc::CursorFetch), // sp_cursorfetch + 8 => Some(RpcProc::CursorOption), // sp_cursoroption (MS-TDS) + 9 => Some(RpcProc::CursorClose), // sp_cursorclose (MS-TDS) + 10 => Some(RpcProc::ExecuteSql), // sp_executesql + 11 => Some(RpcProc::Prepare), // sp_prepare (MS-TDS) + 12 => Some(RpcProc::Execute), // sp_execute (MS-TDS) + 13 => Some(RpcProc::PrepExec), // sp_prepexec + 14 => Some(RpcProc::PrepExecRpc), // sp_prepexecrpc (MS-TDS) + 15 => Some(RpcProc::Unprepare), // sp_unprepare (MS-TDS) + _ => None, + } + } + + pub(crate) fn from_name(name: &str) -> Option { + let n = normalize_proc_name(name); + match n.as_str() { + "sp_cursor" | "sp_cursoropen" => Some(RpcProc::CursorOpen), + "sp_cursorclose" => Some(RpcProc::CursorClose), + "sp_cursorfetch" => Some(RpcProc::CursorFetch), + "sp_cursorprepare" => Some(RpcProc::CursorPrepare), + "sp_cursorexecute" => Some(RpcProc::CursorExecute), + "sp_cursorprepexec" => Some(RpcProc::CursorPrepExec), + "sp_cursorunprepare" => Some(RpcProc::CursorUnprepare), + "sp_cursoroption" => Some(RpcProc::CursorOption), + "sp_executesql" => Some(RpcProc::ExecuteSql), + "sp_prepexec" => Some(RpcProc::PrepExec), + "sp_prepexecrpc" => Some(RpcProc::PrepExecRpc), + "sp_prepare" => Some(RpcProc::Prepare), + "sp_execute" => Some(RpcProc::Execute), + "sp_unprepare" => Some(RpcProc::Unprepare), + "sp_reset_connection" => Some(RpcProc::ResetConnection), + "sp_tables" => Some(RpcProc::SpTables), + "sp_columns" => Some(RpcProc::SpColumns), + "sp_sproc_columns" => Some(RpcProc::SpSprocColumns), + "sp_pkeys" => Some(RpcProc::SpPkeys), + "sp_describe_cursor" => Some(RpcProc::SpDescribeCursor), + _ => None, + } + } + + pub fn is_cursor(&self) -> bool { + matches!( + self, + RpcProc::CursorOpen + | RpcProc::CursorClose + | RpcProc::CursorFetch + | RpcProc::CursorPrepare + | RpcProc::CursorExecute + | RpcProc::CursorPrepExec + | RpcProc::CursorUnprepare + | RpcProc::CursorOption + ) + } +} + +impl CursorOp { + #[allow(dead_code)] + pub(crate) fn from_id(_id: u16) -> Option { + None + } + + #[allow(dead_code)] + pub(crate) fn from_name(name: &str) -> Option { + let n = normalize_proc_name(name); + match n.as_str() { + "sp_cursor" | "sp_cursoropen" => Some(CursorOp::Open), + "sp_cursorclose" => Some(CursorOp::Close), + "sp_cursorfetch" => Some(CursorOp::Fetch), + "sp_cursorprepare" => Some(CursorOp::Prepare), + "sp_cursorexecute" => Some(CursorOp::Execute), + "sp_cursorprepexec" => Some(CursorOp::PrepExec), + "sp_cursorunprepare" => Some(CursorOp::Unprepare), + "sp_cursoroption" => Some(CursorOp::Option), + _ => None, + } + } +} + +#[allow(dead_code)] +pub(crate) fn is_supported_rpc_proc(proc: &RpcProcSelector) -> bool { + match proc { + RpcProcSelector::Id(id) => *id == 10 || *id == 13, + RpcProcSelector::Name(name) => { + let base = normalize_proc_name(name); + base == "sp_executesql" + || base == "sp_prepexec" + || base == "sp_prepare" + || base == "sp_unprepare" + } + } +} + +pub(crate) fn normalize_proc_name(name: &str) -> String { + let mut part = name.trim(); + if let Some(last) = part.rsplit('.').next() { + part = last; + } + part.trim_matches(|c| c == '[' || c == ']' || c == ' ') + .to_ascii_lowercase() +} + +pub fn parse_param_decl(decl: &str) -> Vec<(String, String)> { + if decl.trim().is_empty() { + return vec![]; + } + decl.split(',') + .filter_map(|part| { + let part = part.trim(); + let mut iter = part.splitn(2, char::is_whitespace); + let name = iter.next()?.trim().to_string(); + let type_name = iter.next()?.trim().to_string(); + if name.is_empty() || type_name.is_empty() { + None + } else { + Some((name, type_name)) + } + }) + .collect() +}