use async_trait::async_trait;
use tracing::debug;
use crate::{
connection::execution_context::ExecutionContext,
core::TdsResult,
datatypes::encoder::GenericEncoder,
io::packet_writer::{PacketWriter, TdsPacketWriter},
message::messages::PacketType,
token::tokens::SqlCollation,
};
use super::{
headers::{TdsHeaders, TransactionDescriptorHeader},
messages::Request,
parameters::rpc_parameters::RpcParameter,
};
use crate::message::headers::write_headers;
pub(crate) const PROC_ID_SWITCH: u16 = 0xffff;
pub(crate) const RPC_BATCH_DELIMITER: u8 = 0xff;
#[repr(u8)]
#[derive(Debug, Clone, Copy)]
pub(crate) enum ProcOptions {
None = 0x00,
#[allow(dead_code)]
WithRecompile = 0x01,
#[allow(dead_code)]
NoMetadata = 0x02,
#[allow(dead_code)]
ReuseMetadata = 0x04,
}
pub(crate) enum RpcType {
Named(String),
ProcId(RpcProcs),
}
impl<'a> SqlRpc<'a> {
pub fn new(
rpc_type: RpcType,
positional_parameters: Option<Vec<RpcParameter>>,
named_parameters: Option<Vec<RpcParameter>>,
db_collation: &'a SqlCollation,
execution_context: &ExecutionContext,
) -> Self {
let transaction_descriptor_header: TransactionDescriptorHeader = execution_context.into();
Self {
rpc_type,
headers: Vec::from([transaction_descriptor_header.into()]),
positional_parameters,
named_parameters,
db_collation,
proc_options: ProcOptions::None,
}
}
pub(crate) fn new_batch_command(
rpc_type: RpcType,
positional_parameters: Option<Vec<RpcParameter>>,
named_parameters: Option<Vec<RpcParameter>>,
db_collation: &'a SqlCollation,
execution_context: &ExecutionContext,
first: bool,
) -> Self {
let headers = if first {
Vec::from([TransactionDescriptorHeader::from(execution_context).into()])
} else {
Vec::new()
};
Self {
rpc_type,
headers,
positional_parameters,
named_parameters,
db_collation,
proc_options: ProcOptions::None,
}
}
async fn write_positional_parameters(
&self,
packet_writer: &mut PacketWriter<'_>,
) -> TdsResult<()> {
if let Some(positional_parameters) = &self.positional_parameters {
let encoder = GenericEncoder::new();
for parameter in positional_parameters {
parameter
.serialize(packet_writer, self.db_collation, true, &encoder)
.await?;
}
} else {
debug!("Positional parameters are None, skipping serialization.");
}
Ok(())
}
async fn write_named_parameters(&self, packet_writer: &mut PacketWriter<'_>) -> TdsResult<()> {
if let Some(parameters) = &self.named_parameters {
let encoder = GenericEncoder::new();
for parameter in parameters {
parameter
.serialize(packet_writer, self.db_collation, false, &encoder)
.await?;
}
}
Ok(())
}
pub(crate) async fn serialize_prefix<'s, 'b>(
&'s self,
packet_writer: &'s mut PacketWriter<'b>,
) -> TdsResult<()>
where
'b: 's,
{
write_headers(&self.headers, packet_writer).await?;
self.write_proc(packet_writer).await?;
self.write_positional_parameters(packet_writer).await?;
self.write_named_parameters(packet_writer).await?;
Ok(())
}
pub(crate) async fn serialize_batch_command(
&self,
packet_writer: &mut PacketWriter<'_>,
first: bool,
) -> TdsResult<()> {
if first {
write_headers(&self.headers, packet_writer).await?;
} else {
packet_writer.write_byte_async(RPC_BATCH_DELIMITER).await?;
}
self.write_proc(packet_writer).await?;
self.write_positional_parameters(packet_writer).await?;
self.write_named_parameters(packet_writer).await
}
async fn write_proc(&self, packet_writer: &mut PacketWriter<'_>) -> TdsResult<()> {
match &self.rpc_type {
RpcType::Named(stored_proc_name) => {
packet_writer
.write_i16_async((stored_proc_name.len() as u8).into())
.await?;
packet_writer
.write_string_unicode_async(stored_proc_name.as_str())
.await?;
}
RpcType::ProcId(proc) => {
packet_writer.write_u16_async(PROC_ID_SWITCH).await?;
packet_writer
.write_i16_async(proc.get_u8_value().into())
.await?;
}
}
packet_writer
.write_i16_async(self.proc_options as i16)
.await?;
Ok(())
}
}
#[repr(u8)]
#[derive(Debug, Clone, Copy)]
#[allow(dead_code)]
pub enum RpcProcs {
Cursor = 1,
CursorOpen = 2,
CursorPrepare = 3,
CursorExecute = 4,
CursorPrepExec = 5,
CursorUnprepare = 6,
CursorFetch = 7,
CursorOption = 8,
CursorClose = 9,
ExecuteSql = 10,
Prepare = 11,
Execute = 12,
PrepExec = 13,
PrepExecRpc = 14,
Unprepare = 15,
}
impl RpcProcs {
fn get_u8_value(&self) -> u8 {
*self as u8
}
}
pub(crate) struct SqlRpc<'param> {
pub headers: Vec<TdsHeaders>,
pub rpc_type: RpcType,
pub positional_parameters: Option<Vec<RpcParameter>>,
pub named_parameters: Option<Vec<RpcParameter>>,
pub db_collation: &'param SqlCollation,
pub proc_options: ProcOptions,
}
#[async_trait]
impl Request for SqlRpc<'_> {
fn packet_type(&self) -> PacketType {
PacketType::RpcRequest
}
async fn serialize<'a, 'b>(&'a self, packet_writer: &'a mut PacketWriter<'b>) -> TdsResult<()>
where
'b: 'a,
{
self.serialize_prefix(packet_writer).await?;
packet_writer.finalize().await?;
Ok(())
}
}