pub use crate::tds::stream::{CommandItem, QueryItem, ResultMetadata};
use crate::{
client::Connection,
error::Error,
tds::stream::{CommandReturnValue, ReceivedToken, TokenStream},
FromSql, Row,
};
use futures_util::io::{AsyncRead, AsyncWrite};
use futures_util::stream::TryStreamExt;
use std::fmt::Debug;
#[derive(Debug)]
pub struct ExecuteResult {
rows_affected: Vec<u64>,
}
impl<'a> ExecuteResult {
pub(crate) async fn new<S: AsyncRead + AsyncWrite + Unpin + Send>(
connection: &'a mut Connection<S>,
) -> crate::Result<Self> {
let mut token_stream = TokenStream::new(connection).try_unfold();
let mut rows_affected = Vec::new();
while let Some(token) = token_stream.try_next().await? {
match token {
ReceivedToken::DoneProc(done) if done.is_final() => (),
ReceivedToken::DoneProc(done) => rows_affected.push(done.rows()),
ReceivedToken::DoneInProc(done) => rows_affected.push(done.rows()),
ReceivedToken::Done(done) => rows_affected.push(done.rows()),
_ => (),
}
}
Ok(Self { rows_affected })
}
pub fn rows_affected(&self) -> &[u64] {
self.rows_affected.as_slice()
}
pub fn total(self) -> u64 {
self.rows_affected.into_iter().sum()
}
}
impl IntoIterator for ExecuteResult {
type Item = u64;
type IntoIter = std::vec::IntoIter<Self::Item>;
fn into_iter(self) -> Self::IntoIter {
self.rows_affected.into_iter()
}
}
#[derive(Debug)]
pub struct CommandResult {
pub(crate) rows_affected: Vec<u64>,
pub(crate) return_code: u32,
pub(crate) return_values: Vec<CommandReturnValue>,
pub(crate) query_results: Vec<Vec<Row>>,
}
impl<'a> CommandResult {
pub fn rows_affected(&self) -> &[u64] {
self.rows_affected.as_slice()
}
pub fn return_code(&self) -> u32 {
self.return_code
}
pub fn return_values_len(&self) -> usize {
self.return_values.len()
}
pub fn try_return_value<T>(&'a self, name: &str) -> crate::Result<Option<T>>
where
T: FromSql<'a>,
{
let col_data = self
.return_values
.iter()
.find(|p| p.name.eq(name))
.ok_or_else(|| {
Error::Conversion(format!("Could not find return value {}", name).into())
})?;
T::from_sql(&col_data.data)
}
pub fn to_query_result(&self, idx: usize) -> Option<&Vec<Row>> {
self.query_results.get(idx)
}
}
impl IntoIterator for CommandResult {
type Item = Vec<Row>;
type IntoIter = std::vec::IntoIter<Self::Item>;
fn into_iter(self) -> Self::IntoIter {
self.query_results.into_iter()
}
}
#[cfg(test)]
mod tests {
use super::*;
use crate::tds::codec::ColumnData;
impl ExecuteResult {
fn from_counts(counts: Vec<u64>) -> Self {
Self {
rows_affected: counts,
}
}
}
#[test]
fn execute_result_rows_affected_preserves_order_and_values() {
let res = ExecuteResult::from_counts(vec![3, 0, 7]);
assert_eq!(res.rows_affected(), &[3, 0, 7]);
}
#[test]
fn execute_result_total_sums_every_count() {
assert_eq!(ExecuteResult::from_counts(vec![3, 0, 7]).total(), 10);
}
#[test]
fn execute_result_into_iter_yields_each_count() {
let counts: Vec<u64> = ExecuteResult::from_counts(vec![5, 9]).into_iter().collect();
assert_eq!(counts, vec![5, 9]);
}
fn return_value(name: &str, value: i32) -> CommandReturnValue {
CommandReturnValue {
name: name.to_string(),
ord: 0,
data: ColumnData::I32(Some(value)),
}
}
fn command_result() -> CommandResult {
CommandResult {
rows_affected: vec![2, 4],
return_code: 7,
return_values: vec![return_value("@a", 1), return_value("@b", 42)],
query_results: vec![Vec::new(), Vec::new()],
}
}
#[test]
fn command_result_scalar_accessors() {
let res = command_result();
assert_eq!(res.rows_affected(), &[2, 4]);
assert_eq!(res.return_code(), 7);
assert_eq!(res.return_values_len(), 2);
}
#[test]
fn command_result_to_query_result_indexes_record_sets() {
let res = command_result();
assert!(res.to_query_result(0).is_some());
assert!(res.to_query_result(1).is_some());
assert!(res.to_query_result(2).is_none());
}
#[test]
fn command_result_try_return_value_reads_named_out_param() {
let res = command_result();
let got: Option<i32> = res.try_return_value("@b").unwrap();
assert_eq!(got, Some(42));
assert!(res.try_return_value::<i32>("@missing").is_err());
}
#[test]
fn command_result_into_iter_yields_each_record_set() {
assert_eq!(command_result().into_iter().count(), 2);
}
}