#![forbid(unsafe_code)]
use std::collections::{BTreeMap, HashMap, HashSet, VecDeque};
use std::process;
use std::sync::atomic::{AtomicBool, Ordering};
use std::sync::Arc;
use std::time::Duration;
use asupersync::io::{AsyncRead, AsyncReadExt, AsyncWrite, AsyncWriteExt};
use asupersync::net::TcpStream;
use asupersync::runtime::{reactor, Runtime, RuntimeBuilder};
use asupersync::sync::Mutex as AsyncMutex;
use asupersync::{time, Cx};
use oracledb_protocol::thin::aq::{
build_aq_array_deq_payload, build_aq_array_enq_payload, build_aq_deq_payload,
build_aq_enq_payload, parse_aq_array_response, parse_aq_deq_response, parse_aq_enq_response,
AqArrayResult, AqDeqOptions, AqDeqResult, AqEnqOptions, AqMsgProps, AqQueueDesc,
};
use oracledb_protocol::thin::{
adjust_refetch_metadata, build_auth_phase_two_payload_with_proxy_with_seq,
build_begin_pipeline_piggyback, build_change_password_payload_with_seq,
build_connect_packet_payload, build_define_fetch_payload_with_seq,
build_end_pipeline_payload_with_seq, build_execute_payload_with_bind_rows_and_options_with_seq,
build_execute_payload_with_bind_rows_with_seq_and_token, build_execute_payload_with_seq,
build_fast_auth_phase_one_payload, build_fetch_payload_with_seq,
build_function_payload_with_seq, build_function_payload_with_seq_and_token,
build_lob_create_temp_payload_with_seq, build_lob_free_temp_payload_with_seq,
build_lob_read_payload_with_seq, build_lob_trim_payload_with_seq,
build_lob_write_payload_with_seq, parse_accept_payload, parse_auth_response,
parse_fetch_response_with_context, parse_lob_create_temp_response,
parse_lob_free_temp_response, parse_lob_read_response, parse_lob_trim_response,
parse_lob_write_response, parse_plain_function_response, parse_query_response,
parse_query_response_borrowed, parse_query_response_with_binds_options_and_columns,
parse_tpc_txn_switch_response, BindValue, BorrowedFetchResult, ClientCapabilities,
ColumnMetadata, ExecuteOptions, LobReadResult, QueryResult, QueryValueRef, SessionlessTxnState,
TpcChangeStateResponse, TpcSwitchResponse, TpcXid, TNS_DATA_FLAGS_BEGIN_PIPELINE,
TNS_DATA_FLAGS_END_OF_REQUEST, TNS_FUNC_COMMIT, TNS_FUNC_LOGOFF, TNS_FUNC_PING,
TNS_FUNC_ROLLBACK, TNS_MSG_TYPE_END_OF_RESPONSE, TNS_MSG_TYPE_FLUSH_OUT_BINDS,
TNS_PACKET_TYPE_ACCEPT, TNS_PACKET_TYPE_CONNECT, TNS_PACKET_TYPE_DATA,
TNS_PACKET_TYPE_REDIRECT, TNS_PACKET_TYPE_REFUSE, TNS_PIPELINE_MODE_ABORT_ON_ERROR,
TNS_PIPELINE_MODE_CONTINUE_ON_ERROR, TNS_TPC_TXN_ABORT, TNS_TPC_TXN_COMMIT, TNS_TPC_TXN_DETACH,
TNS_TPC_TXN_POST_DETACH, TNS_TPC_TXN_PREPARE, TNS_TPC_TXN_START, TNS_TPC_TXN_STATE_ABORTED,
TNS_TPC_TXN_STATE_COMMITTED, TNS_TPC_TXN_STATE_FORGOTTEN, TNS_TPC_TXN_STATE_PREPARE,
TNS_TPC_TXN_STATE_READ_ONLY, TNS_TPC_TXN_STATE_REQUIRES_COMMIT, TPC_TXN_FLAGS_NEW,
TPC_TXN_FLAGS_RESUME, TPC_TXN_FLAGS_SESSIONLESS,
};
use oracledb_protocol::thin::{
build_notify_payload_with_seq, build_subscribe_payload_with_seq, check_notification_header,
parse_subscribe_response, try_parse_oac_record, NotificationRecord, SubscribeResult,
TNS_SUBSCR_OP_REGISTER, TNS_SUBSCR_OP_UNREGISTER,
};
use oracledb_protocol::thin::{
build_sessionless_piggyback, build_tpc_change_state_payload_with_seq,
build_tpc_switch_payload_with_seq, build_tpc_txn_switch_payload_with_seq,
parse_tpc_change_state_response, parse_tpc_switch_response,
};
use oracledb_protocol::thin::{TNS_AQ_ARRAY_DEQ, TNS_AQ_ARRAY_ENQ};
use oracledb_protocol::wire::{encode_packet, PacketLengthWidth};
use oracledb_protocol::{net::EasyConnect, ClientIdentity};
const PYTHON_ORACLEDB_COMPAT_VERSION_NUM: u32 = 0x0400_1000;
const DEFAULT_SDU: usize = 8192;
const TNS_DATA_PACKET_OVERHEAD: usize = 10;
pub use oracledb_protocol as protocol;
mod fetch_profile {
use std::sync::atomic::{AtomicBool, AtomicU64, Ordering};
use std::sync::OnceLock;
static READ_NS: AtomicU64 = AtomicU64::new(0);
static DECODE_NS: AtomicU64 = AtomicU64::new(0);
static ENABLED: OnceLock<bool> = OnceLock::new();
static FORCE: AtomicBool = AtomicBool::new(false);
#[inline]
pub(crate) fn enabled() -> bool {
FORCE.load(Ordering::Relaxed)
|| *ENABLED.get_or_init(|| std::env::var_os("ORACLEDB_PROFILE_FETCH").is_some())
}
#[inline]
pub(crate) fn add_read(ns: u64) {
READ_NS.fetch_add(ns, Ordering::Relaxed);
}
#[inline]
pub(crate) fn add_decode(ns: u64) {
DECODE_NS.fetch_add(ns, Ordering::Relaxed);
}
pub(crate) fn snapshot() -> (u64, u64) {
(
READ_NS.load(Ordering::Relaxed),
DECODE_NS.load(Ordering::Relaxed),
)
}
pub(crate) fn reset() {
READ_NS.store(0, Ordering::Relaxed);
DECODE_NS.store(0, Ordering::Relaxed);
}
pub(crate) fn set_force(on: bool) {
FORCE.store(on, Ordering::Relaxed);
}
}
pub fn fetch_profile_read_decode_ns() -> (u64, u64) {
fetch_profile::snapshot()
}
pub fn fetch_profile_reset() {
fetch_profile::reset();
}
pub fn fetch_profile_arm(on: bool) {
fetch_profile::set_force(on);
}
#[cfg(feature = "arrow")]
pub mod arrow;
pub mod cursor_logic;
#[macro_use]
mod obs;
pub mod pool;
#[cfg(feature = "soda")]
pub mod soda;
mod sql_convert;
pub mod tls;
pub mod transport;
#[cfg(feature = "tracing")]
#[doc(hidden)]
pub use tracing as __tracing;
#[cfg(not(feature = "tracing"))]
#[doc(hidden)]
pub use obs::ObsSpanGuard;
pub use cursor_logic::{
bind_rows_need_iterative_plsql, ExecutemanyManager, ExecutemanyManagerError,
};
pub use sql_convert::{
ConversionError, FromRow, FromSql, IntoBinds, QueryResultExt, ToSql, TypedRow,
};
#[cfg(feature = "derive")]
pub use oracledb_derive::FromRow;
use transport::{OracleReadHalf, OracleWriteHalf};
type SharedWriteHalf = Arc<AsyncMutex<OracleWriteHalf>>;
const SESSION_DEAD_ORA_CODES: &[u32] = &[
22, 28, 31, 45, 378, 600, 602, 603, 609, 1012, 1041, 1043, 1089, 1092, 2396, 3113, 3114, 3122,
3135, 12153, 12537, 12547, 12570, 12583, 27146, 28511, 56600,
];
const TNS_CCAP_FIELD_VERSION_18_1_EXT_1: u8 = 11;
fn decode_server_version_number(full: u32, new_format: bool) -> (u8, u8, u8, u8, u8) {
if new_format {
(
((full >> 24) & 0xFF) as u8,
((full >> 16) & 0xFF) as u8,
((full >> 12) & 0x0F) as u8,
((full >> 4) & 0xFF) as u8,
(full & 0x0F) as u8,
)
} else {
(
((full >> 24) & 0xFF) as u8,
((full >> 20) & 0x0F) as u8,
((full >> 12) & 0x0F) as u8,
((full >> 8) & 0x0F) as u8,
(full & 0x0F) as u8,
)
}
}
const TRANSIENT_ORA_CODES: &[u32] = &[54, 60, 104, 257, 12516, 12520, 12526, 12528, 30006, 51535];
const CONNECTION_LOST_ORA_CODES: &[u32] = &[
28, 1012, 1041, 1089, 2396, 3113, 3114, 3135, 12537, 12547, 12570, 28511,
];
fn parse_ora_code_from_message(message: &str) -> Option<u32> {
let start = message.find("ORA-")?;
let digits: String = message[start + 4..]
.chars()
.take_while(|ch| ch.is_ascii_digit())
.collect();
digits.parse::<u32>().ok()
}
fn protocol_error_is_session_dead(err: &oracledb_protocol::ProtocolError) -> bool {
protocol_error_ora_code(err).is_some_and(|code| SESSION_DEAD_ORA_CODES.contains(&code))
}
fn protocol_error_ora_code(err: &oracledb_protocol::ProtocolError) -> Option<u32> {
match err {
oracledb_protocol::ProtocolError::ServerError(message) => {
parse_ora_code_from_message(message)
}
oracledb_protocol::ProtocolError::ServerErrorWithRowCount { message, .. } => {
parse_ora_code_from_message(message)
}
oracledb_protocol::ProtocolError::ServerErrorInfo(details) => Some(details.code),
_ => None,
}
}
fn protocol_error_offset(err: &oracledb_protocol::ProtocolError) -> Option<i32> {
match err {
oracledb_protocol::ProtocolError::ServerErrorInfo(details) if details.pos != 0 => {
Some(details.pos)
}
_ => None,
}
}
#[derive(Debug, thiserror::Error)]
pub enum Error {
#[error(transparent)]
Protocol(#[from] oracledb_protocol::ProtocolError),
#[error("I/O error: {0}")]
Io(#[from] std::io::Error),
#[error("asupersync runtime error: {0}")]
Runtime(String),
#[error("listener redirected this connection; redirect handling is not implemented yet")]
RedirectUnsupported,
#[error("listener refused connection: {0}")]
ListenerRefused(String),
#[error("server did not advertise fast authentication")]
FastAuthRequired,
#[error("server response did not contain {0}")]
MissingSessionField(&'static str),
#[error("call timeout of {0} ms exceeded")]
CallTimeout(u32),
#[error("ORA-01013: user requested cancel of current operation")]
Cancelled,
#[error("DPY-4011: the database or network closed the connection: {0}")]
ConnectionClosed(String),
#[error("TLS/TCPS error: {0}")]
Tls(String),
#[error("{0}")]
SessionlessTransaction(SessionlessError),
#[error("DPY-5010: internal error: unknown transaction state {0}")]
UnknownTransactionState(u32),
#[error("type conversion failed: {0}")]
Conversion(ConversionError),
#[cfg(feature = "arrow")]
#[error(transparent)]
ArrowConversion(#[from] arrow::ArrowConversionError),
}
pub type Result<T> = std::result::Result<T, Error>;
impl Error {
pub fn ora_code(&self) -> Option<i32> {
match self {
Error::Protocol(err) => protocol_error_ora_code(err).map(|code| code as i32),
Error::Cancelled => Some(1013),
_ => None,
}
}
pub fn offset(&self) -> Option<i32> {
match self {
Error::Protocol(err) => protocol_error_offset(err),
_ => None,
}
}
pub fn is_connection_lost(&self) -> bool {
match self {
Error::Io(_) | Error::ConnectionClosed(_) => true,
_ => self
.ora_code()
.is_some_and(|code| CONNECTION_LOST_ORA_CODES.contains(&(code as u32))),
}
}
pub fn is_transient(&self) -> bool {
matches!(self, Error::CallTimeout(_) | Error::Cancelled)
|| self
.ora_code()
.is_some_and(|code| TRANSIENT_ORA_CODES.contains(&(code as u32)))
}
pub fn is_retryable(&self) -> bool {
self.is_transient() || self.is_connection_lost()
}
}
#[derive(Clone, Copy, Debug, Eq, PartialEq)]
pub enum SessionlessError {
DifferingMethods,
AlreadyActive,
Inactive,
}
impl SessionlessError {
pub fn full_code(self) -> &'static str {
match self {
Self::DifferingMethods => "DPY-3034",
Self::AlreadyActive => "DPY-3035",
Self::Inactive => "DPY-3036",
}
}
pub fn message(self) -> &'static str {
match self {
Self::DifferingMethods => {
"suspending or resuming a Sessionless Transaction can be done with \
DBMS_TRANSACTION or with python-oracledb, but not both"
}
Self::AlreadyActive => {
"suspend, commit, or rollback the current active sessionless \
transaction before beginning or resuming another one"
}
Self::Inactive => "no Sessionless Transaction is active",
}
}
}
impl std::fmt::Display for SessionlessError {
fn fmt(&self, f: &mut std::fmt::Formatter<'_>) -> std::fmt::Result {
write!(f, "{}: {}", self.full_code(), self.message())
}
}
fn refetch_retry_applies(err: &Error) -> bool {
let message = match err {
Error::Protocol(oracledb_protocol::ProtocolError::ServerError(message)) => message,
Error::Protocol(oracledb_protocol::ProtocolError::ServerErrorWithRowCount {
message,
..
}) => message,
Error::Protocol(oracledb_protocol::ProtocolError::ServerErrorInfo(details)) => {
return details.code == 932 || details.code == 1007;
}
_ => return false,
};
message.starts_with("ORA-00932") || message.starts_with("ORA-01007")
}
#[derive(Clone, Debug)]
pub struct ConnectOptions {
pub connect_string: String,
pub user: String,
pub password: String,
pub identity: ClientIdentity,
pub app_context: Vec<(String, String, String)>,
pub sdu: u16,
pub proxy_user: Option<String>,
pub server_type_emon: bool,
pub wallet_location: Option<String>,
pub wallet_password: Option<String>,
pub ssl_server_dn_match: bool,
pub ssl_server_cert_dn: Option<String>,
pub use_sni: bool,
}
impl ConnectOptions {
pub fn new(
connect_string: impl Into<String>,
user: impl Into<String>,
password: impl Into<String>,
identity: ClientIdentity,
) -> Self {
Self {
connect_string: connect_string.into(),
user: user.into(),
password: password.into(),
identity,
app_context: Vec::new(),
sdu: 8192,
proxy_user: None,
server_type_emon: false,
wallet_location: None,
wallet_password: None,
ssl_server_dn_match: true,
ssl_server_cert_dn: None,
use_sni: false,
}
}
#[must_use]
pub fn with_use_sni(mut self, use_sni: bool) -> Self {
self.use_sni = use_sni;
self
}
#[must_use]
pub fn with_wallet_location(mut self, location: impl Into<String>) -> Self {
self.wallet_location = Some(location.into());
self
}
#[must_use]
pub fn with_wallet_password(mut self, password: impl Into<String>) -> Self {
self.wallet_password = Some(password.into());
self
}
#[must_use]
pub fn with_ssl_server_dn_match(mut self, enabled: bool) -> Self {
self.ssl_server_dn_match = enabled;
self
}
#[must_use]
pub fn with_ssl_server_cert_dn(mut self, dn: impl Into<String>) -> Self {
self.ssl_server_cert_dn = Some(dn.into());
self
}
pub fn with_server_type_emon(mut self, emon: bool) -> Self {
self.server_type_emon = emon;
self
}
pub fn with_app_context(mut self, app_context: Vec<(String, String, String)>) -> Self {
self.app_context = app_context;
self
}
pub fn with_proxy_user(mut self, proxy_user: Option<String>) -> Self {
self.proxy_user = proxy_user;
self
}
pub fn with_sdu(mut self, sdu: u32) -> Self {
let clamped = sdu.clamp(512, u32::from(u16::MAX));
self.sdu = u16::try_from(clamped).unwrap_or(u16::MAX);
self
}
}
#[derive(Debug)]
pub struct Connection {
descriptor: EasyConnect,
identity: ClientIdentity,
read: OracleReadHalf,
write: SharedWriteHalf,
session_id: u32,
serial_num: u16,
server_version: Option<String>,
server_version_tuple: Option<(u8, u8, u8, u8, u8)>,
capabilities: ClientCapabilities,
ttc_seq_num: u8,
sdu: usize,
supports_end_of_response: bool,
supports_oob: bool,
cancel_drain_pending: Arc<AtomicBool>,
cursor_columns: BTreeMap<u32, Vec<ColumnMetadata>>,
fetch_metadata_by_sql: HashMap<String, Vec<ColumnMetadata>>,
fetch_metadata_order: VecDeque<String>,
dead: bool,
user: String,
combo_key: Vec<u8>,
statement_cache: Vec<(String, u32)>,
in_use_cursors: HashSet<u32>,
copied_cursors: HashSet<u32>,
cursors_to_close: Vec<u32>,
sessionless_data: Option<SessionlessData>,
notification_buffer: Vec<u8>,
notification_header_consumed: bool,
transaction_context: Option<Vec<u8>>,
txn_in_progress: bool,
}
#[derive(Clone, Debug)]
struct SessionlessData {
transaction_id: Vec<u8>,
timeout: u32,
operation: u32,
flags: u32,
piggyback_pending: bool,
started_on_server: bool,
}
enum PacketRead {
Appended,
TimedOut,
Closed,
}
#[derive(Clone, Debug)]
pub enum NotificationOutcome {
Record(NotificationRecord),
TimedOut,
Closed,
}
const STATEMENT_CACHE_SIZE: usize = 20;
#[derive(Clone, Debug)]
pub enum PipelineRequest {
Execute {
sql: String,
bind_rows: Vec<Vec<BindValue>>,
prefetch_rows: u32,
},
Commit,
}
#[derive(Debug)]
pub struct CancelHandle {
write: SharedWriteHalf,
}
impl Connection {
pub async fn connect(cx: &Cx, options: ConnectOptions) -> Result<Self> {
cx.checkpoint()
.map_err(|err| Error::Runtime(err.to_string()))?;
let descriptor = EasyConnect::parse(&options.connect_string)?;
let _span = obs_span!(
"oracledb.connect",
db.system = "oracle",
server.address = %descriptor.host,
server.port = descriptor.port as u64,
db.name = %descriptor.service_name,
);
let identity = options.identity;
trace_connect_step("tcp connect");
let stream = TcpStream::connect_timeout(
(descriptor.host.clone(), descriptor.port),
Duration::from_secs(20),
)
.await?;
stream.set_nodelay(true)?;
trace_connect_step("tcp connected");
let (mut read, write) = if descriptor.protocol.is_tls() {
trace_connect_step("tls handshake");
let server_type = if options.server_type_emon {
Some("emon")
} else {
None
};
let tls_params = tls::resolve_tls_params(
&descriptor,
options.wallet_location.as_deref(),
options.wallet_password.as_deref(),
options.ssl_server_dn_match,
options.ssl_server_cert_dn.as_deref(),
options.use_sni,
)?;
let tls_stream =
tls::tls_handshake(&descriptor, server_type, &tls_params, stream).await?;
trace_connect_step("tls established");
transport::tls_split(tls_stream)
} else {
transport::plain_split(stream)
};
let write = Arc::new(AsyncMutex::with_name("oracle_tcp_write", write));
let connect_descriptor = listener_connect_descriptor_with_server(
&descriptor,
&identity,
options.server_type_emon,
);
trace_connect_value("CONNECT descriptor", &connect_descriptor);
let connect_payload = build_connect_packet_payload(&connect_descriptor, options.sdu)?;
let packet = encode_packet(
TNS_PACKET_TYPE_CONNECT,
0,
None,
&connect_payload,
PacketLengthWidth::Legacy16,
)?;
trace_connect_bytes("CONNECT packet", &packet);
trace_connect_step("send CONNECT");
write_all_shared(cx, &write, &packet).await?;
trace_connect_step("read ACCEPT");
let accept = read_packet(&mut read, PacketLengthWidth::Legacy16).await?;
match accept.packet_type {
TNS_PACKET_TYPE_ACCEPT => {}
TNS_PACKET_TYPE_REDIRECT => return Err(Error::RedirectUnsupported),
TNS_PACKET_TYPE_REFUSE => {
return Err(Error::ListenerRefused(
String::from_utf8_lossy(&accept.payload).to_string(),
))
}
other => {
return Err(oracledb_protocol::ProtocolError::UnknownMessageType {
message_type: other,
position: 4,
}
.into())
}
}
let accept_info = parse_accept_payload(&accept.payload)?;
if !accept_info.supports_fast_auth {
return Err(Error::FastAuthRequired);
}
let sdu = usize::try_from(accept_info.sdu)
.unwrap_or(DEFAULT_SDU)
.max(TNS_DATA_PACKET_OVERHEAD + 1);
let client_pid = process::id();
let auth_one = build_fast_auth_phase_one_payload(
&options.user,
&identity.program,
&identity.machine,
&identity.osuser,
&identity.terminal,
client_pid,
)?;
trace_connect_bytes("AUTH phase one payload", &auth_one);
trace_connect_step("send AUTH phase one");
send_data_packet_shared(cx, &write, &auth_one, sdu).await?;
trace_connect_step("read AUTH phase one");
let auth_one_response = read_data_response(&mut read, cx, &write).await?;
trace_connect_bytes("AUTH phase one response", &auth_one_response);
let auth_one = parse_auth_response(&auth_one_response)?;
let capabilities = auth_one.capabilities.unwrap_or_default();
let mut ttc_seq_num = 1;
let verifier_type = auth_one
.verifier_type
.ok_or(Error::MissingSessionField("AUTH_VFR_DATA verifier type"))?;
let encrypted = oracledb_protocol::crypto::generate_verifier(
options.password.as_bytes(),
&auth_one.session_data,
verifier_type,
)?;
let auth_connect_string = auth_connect_descriptor(&descriptor);
let auth_two = build_auth_phase_two_payload_with_proxy_with_seq(
&options.user,
&encrypted,
&identity.driver_name,
PYTHON_ORACLEDB_COMPAT_VERSION_NUM,
&auth_connect_string,
next_ttc_sequence(&mut ttc_seq_num),
&options.app_context,
options.proxy_user.as_deref(),
)?;
trace_connect_bytes("AUTH phase two payload", &auth_two);
trace_connect_step("send AUTH phase two");
send_data_packet_shared(cx, &write, &auth_two, sdu).await?;
trace_connect_step("read AUTH phase two");
let auth_two_response = read_data_response(&mut read, cx, &write).await?;
trace_connect_bytes("AUTH phase two response", &auth_two_response);
let auth_two = parse_auth_response(&auth_two_response)?;
oracledb_protocol::crypto::verify_server_response(
&encrypted.combo_key,
&auth_two.session_data,
)?;
let session_id = parse_session_u32(&auth_two.session_data, "AUTH_SESSION_ID")?;
let serial_num = parse_session_u16(&auth_two.session_data, "AUTH_SERIAL_NUM")?;
let server_version = auth_two.session_data.get("AUTH_VERSION_STRING").cloned();
let server_version_tuple = auth_two
.session_data
.get("AUTH_VERSION_NO")
.and_then(|value| value.trim().parse::<u32>().ok())
.map(|num| {
decode_server_version_number(
num,
capabilities.ttc_field_version >= TNS_CCAP_FIELD_VERSION_18_1_EXT_1,
)
});
Ok(Self {
descriptor,
identity,
read,
write,
session_id,
serial_num,
server_version,
server_version_tuple,
capabilities,
ttc_seq_num,
sdu,
supports_end_of_response: accept_info.supports_end_of_response,
supports_oob: accept_info.supports_oob,
cancel_drain_pending: Arc::new(AtomicBool::new(false)),
cursor_columns: BTreeMap::new(),
fetch_metadata_by_sql: HashMap::new(),
fetch_metadata_order: VecDeque::new(),
dead: false,
user: options.user,
combo_key: encrypted.combo_key,
statement_cache: Vec::new(),
in_use_cursors: HashSet::new(),
copied_cursors: HashSet::new(),
cursors_to_close: Vec::new(),
sessionless_data: None,
notification_buffer: Vec::new(),
notification_header_consumed: false,
transaction_context: None,
txn_in_progress: false,
})
}
pub fn descriptor(&self) -> &EasyConnect {
&self.descriptor
}
pub fn identity(&self) -> &ClientIdentity {
&self.identity
}
pub fn session_id(&self) -> u32 {
self.session_id
}
pub fn serial_num(&self) -> u16 {
self.serial_num
}
pub fn server_version(&self) -> Option<&str> {
self.server_version.as_deref()
}
pub fn server_version_tuple(&self) -> Option<(u8, u8, u8, u8, u8)> {
self.server_version_tuple
}
fn supports_oson_long_fnames(&self) -> bool {
self.server_version_tuple
.map(|(major, ..)| major >= 23)
.unwrap_or(false)
}
pub fn sdu(&self) -> usize {
self.sdu
}
pub fn supports_pipelining(&self) -> bool {
self.supports_end_of_response
}
pub fn cancel_handle(&self) -> Result<CancelHandle> {
Ok(CancelHandle {
write: Arc::clone(&self.write),
})
}
pub fn is_dead(&self) -> bool {
self.dead
}
fn note_parse<T>(
&mut self,
result: std::result::Result<T, oracledb_protocol::ProtocolError>,
) -> Result<T> {
match result {
Ok(value) => Ok(value),
Err(err) => {
if protocol_error_is_session_dead(&err) {
self.dead = true;
}
Err(Error::Protocol(err))
}
}
}
pub async fn ping(&mut self, cx: &Cx) -> Result<()> {
self.send_function(cx, TNS_FUNC_PING).await
}
pub async fn change_password(
&mut self,
cx: &Cx,
old_password: &str,
new_password: &str,
) -> Result<()> {
cx.checkpoint()
.map_err(|err| Error::Runtime(err.to_string()))?;
let (encoded_password, encoded_newpassword) =
oracledb_protocol::crypto::encrypt_change_password_pair(
&self.combo_key,
old_password.as_bytes(),
new_password.as_bytes(),
)?;
let seq_num = next_ttc_sequence(&mut self.ttc_seq_num);
let payload = build_change_password_payload_with_seq(
&self.user,
&encoded_password,
&encoded_newpassword,
seq_num,
)?;
send_data_packet_shared(cx, &self.write, &payload, self.sdu).await?;
let response = read_data_response(&mut self.read, cx, &self.write).await?;
self.note_parse(parse_auth_response(&response).map(|_| ()))?;
Ok(())
}
#[allow(clippy::too_many_arguments)]
pub async fn subscribe_register(
&mut self,
cx: &Cx,
namespace: u32,
name: Option<&str>,
public_qos: u32,
operations: u32,
timeout: u32,
grouping_class: u8,
grouping_value: u32,
grouping_type: u8,
) -> Result<SubscribeResult> {
cx.checkpoint()
.map_err(|err| Error::Runtime(err.to_string()))?;
let seq_num = next_ttc_sequence(&mut self.ttc_seq_num);
let payload = build_subscribe_payload_with_seq(
seq_num,
TNS_SUBSCR_OP_REGISTER,
Some(&self.user),
None,
namespace,
name,
public_qos,
operations,
timeout,
grouping_class,
grouping_value,
grouping_type,
0,
self.capabilities.ttc_field_version,
)?;
send_data_packet_shared(cx, &self.write, &payload, self.sdu).await?;
let response = read_data_response(&mut self.read, cx, &self.write).await?;
self.note_parse(parse_subscribe_response(&response, self.capabilities))
}
#[allow(clippy::too_many_arguments)]
pub async fn subscribe_unregister(
&mut self,
cx: &Cx,
registration_id: u64,
client_id: &[u8],
namespace: u32,
name: Option<&str>,
public_qos: u32,
operations: u32,
timeout: u32,
grouping_class: u8,
grouping_value: u32,
grouping_type: u8,
) -> Result<()> {
cx.checkpoint()
.map_err(|err| Error::Runtime(err.to_string()))?;
let seq_num = next_ttc_sequence(&mut self.ttc_seq_num);
let payload = build_subscribe_payload_with_seq(
seq_num,
TNS_SUBSCR_OP_UNREGISTER,
Some(&self.user),
Some(client_id),
namespace,
name,
public_qos,
operations,
timeout,
grouping_class,
grouping_value,
grouping_type,
registration_id,
self.capabilities.ttc_field_version,
)?;
send_data_packet_shared(cx, &self.write, &payload, self.sdu).await?;
let response = read_data_response(&mut self.read, cx, &self.write).await?;
self.note_parse(parse_subscribe_response(&response, self.capabilities))?;
Ok(())
}
pub async fn notify_register(&mut self, cx: &Cx, client_id: &[u8]) -> Result<()> {
cx.checkpoint()
.map_err(|err| Error::Runtime(err.to_string()))?;
let seq_num = next_ttc_sequence(&mut self.ttc_seq_num);
let payload =
build_notify_payload_with_seq(seq_num, client_id, self.capabilities.ttc_field_version)?;
send_data_packet_shared_with_flags(
cx,
&self.write,
&payload,
self.sdu,
0,
TNS_DATA_FLAGS_END_OF_REQUEST,
)
.await?;
Ok(())
}
pub async fn recv_notification(
&mut self,
cx: &Cx,
namespace: u32,
public_qos: u32,
read_timeout: Duration,
) -> Result<NotificationOutcome> {
cx.checkpoint()
.map_err(|err| Error::Runtime(err.to_string()))?;
let db_name = self.descriptor.service_name.clone();
loop {
if !self.notification_header_consumed {
if self.notification_buffer.is_empty() {
match self.read_one_notification_packet(read_timeout).await? {
PacketRead::Appended => continue,
PacketRead::TimedOut => return Ok(NotificationOutcome::TimedOut),
PacketRead::Closed => return Ok(NotificationOutcome::Closed),
}
}
let consumed = check_notification_header(&self.notification_buffer)?;
self.notification_buffer.drain(..consumed);
self.notification_header_consumed = true;
}
if !self.notification_buffer.is_empty() {
if let Some((record, consumed)) = try_parse_oac_record(
&self.notification_buffer,
namespace,
public_qos,
Some(&db_name),
)? {
self.notification_buffer.drain(..consumed);
return Ok(NotificationOutcome::Record(record));
}
}
match self.read_one_notification_packet(read_timeout).await? {
PacketRead::Appended => {}
PacketRead::TimedOut => return Ok(NotificationOutcome::TimedOut),
PacketRead::Closed => return Ok(NotificationOutcome::Closed),
}
}
}
pub async fn execute_query_for_registration(
&mut self,
cx: &Cx,
sql: &str,
registration_id: u64,
) -> Result<Option<u64>> {
let exec_options = ExecuteOptions {
registration_id,
..ExecuteOptions::default()
};
let result = self
.execute_query_with_bind_rows_and_options(cx, sql, 0, &[], exec_options)
.await?;
Ok(result.query_id)
}
async fn read_one_notification_packet(&mut self, read_timeout: Duration) -> Result<PacketRead> {
let read = read_packet(&mut self.read, PacketLengthWidth::Large32);
let packet = match time::timeout(time::wall_now(), read_timeout, read).await {
Ok(Ok(packet)) => packet,
Ok(Err(_)) => return Ok(PacketRead::Closed),
Err(_) => return Ok(PacketRead::TimedOut),
};
if packet.packet_type != TNS_PACKET_TYPE_DATA {
return Ok(PacketRead::Closed);
}
let Some((_data_flags, payload)) = packet.payload.split_at_checked(2) else {
return Ok(PacketRead::Closed);
};
self.notification_buffer.extend_from_slice(payload);
Ok(PacketRead::Appended)
}
pub async fn ping_with_timeout(&mut self, cx: &Cx, timeout_ms: u32) -> Result<()> {
if timeout_ms == 0 {
return self.ping(cx).await;
}
match time::timeout(
time::wall_now(),
Duration::from_millis(u64::from(timeout_ms)),
self.ping(cx),
)
.await
{
Ok(result) => result,
Err(_) => self.recover_from_call_timeout(cx, timeout_ms).await,
}
}
pub async fn commit(&mut self, cx: &Cx) -> Result<()> {
let _span = obs_span!("oracledb.commit");
self.send_function(cx, TNS_FUNC_COMMIT).await?;
self.sessionless_data = None;
Ok(())
}
pub async fn rollback(&mut self, cx: &Cx) -> Result<()> {
let _span = obs_span!("oracledb.rollback");
self.send_function(cx, TNS_FUNC_ROLLBACK).await?;
self.sessionless_data = None;
Ok(())
}
async fn start_sessionless_transaction(
&mut self,
cx: &Cx,
transaction_id: &[u8],
timeout: u32,
flags: u32,
defer_round_trip: bool,
) -> Result<()> {
if self.sessionless_data.is_some() {
return Err(Error::SessionlessTransaction(
SessionlessError::AlreadyActive,
));
}
let data = SessionlessData {
transaction_id: transaction_id.to_vec(),
timeout,
operation: TNS_TPC_TXN_START,
flags,
piggyback_pending: defer_round_trip,
started_on_server: false,
};
if defer_round_trip {
self.sessionless_data = Some(data);
return Ok(());
}
cx.checkpoint()
.map_err(|err| Error::Runtime(err.to_string()))?;
let seq_num = next_ttc_sequence(&mut self.ttc_seq_num);
let payload = build_tpc_txn_switch_payload_with_seq(
seq_num,
0,
data.operation,
data.flags | TPC_TXN_FLAGS_SESSIONLESS,
data.timeout,
Some(transaction_id),
);
send_data_packet_shared(cx, &self.write, &payload, self.sdu).await?;
let response = read_data_response(&mut self.read, cx, &self.write).await?;
let state = self.note_parse(parse_tpc_txn_switch_response(&response, self.capabilities))?;
self.sessionless_data = Some(data);
self.apply_sessionless_state(state);
Ok(())
}
pub async fn begin_sessionless_transaction(
&mut self,
cx: &Cx,
transaction_id: &[u8],
timeout: u32,
defer_round_trip: bool,
) -> Result<()> {
self.start_sessionless_transaction(
cx,
transaction_id,
timeout,
TPC_TXN_FLAGS_NEW,
defer_round_trip,
)
.await
}
pub async fn resume_sessionless_transaction(
&mut self,
cx: &Cx,
transaction_id: &[u8],
timeout: u32,
defer_round_trip: bool,
) -> Result<()> {
self.start_sessionless_transaction(
cx,
transaction_id,
timeout,
TPC_TXN_FLAGS_RESUME,
defer_round_trip,
)
.await
}
pub async fn suspend_sessionless_transaction(&mut self, cx: &Cx) -> Result<()> {
match &self.sessionless_data {
None => return Err(Error::SessionlessTransaction(SessionlessError::Inactive)),
Some(data) if data.started_on_server => {
return Err(Error::SessionlessTransaction(
SessionlessError::DifferingMethods,
));
}
Some(_) => {}
}
cx.checkpoint()
.map_err(|err| Error::Runtime(err.to_string()))?;
let seq_num = next_ttc_sequence(&mut self.ttc_seq_num);
let payload = build_tpc_txn_switch_payload_with_seq(
seq_num,
0,
TNS_TPC_TXN_DETACH,
TPC_TXN_FLAGS_SESSIONLESS,
0,
None,
);
send_data_packet_shared(cx, &self.write, &payload, self.sdu).await?;
let response = read_data_response(&mut self.read, cx, &self.write).await?;
let state = self.note_parse(parse_tpc_txn_switch_response(&response, self.capabilities))?;
self.sessionless_data = None;
self.apply_sessionless_state(state);
Ok(())
}
async fn tpc_switch_round_trip(
&mut self,
cx: &Cx,
operation: u32,
flags: u32,
timeout: u32,
xid: Option<&TpcXid<'_>>,
context: Option<&[u8]>,
) -> Result<TpcSwitchResponse> {
cx.checkpoint()
.map_err(|err| Error::Runtime(err.to_string()))?;
let seq_num = next_ttc_sequence(&mut self.ttc_seq_num);
let payload =
build_tpc_switch_payload_with_seq(seq_num, operation, flags, timeout, xid, context);
send_data_packet_shared(cx, &self.write, &payload, self.sdu).await?;
let response = read_data_response(&mut self.read, cx, &self.write).await?;
self.note_parse(parse_tpc_switch_response(&response, self.capabilities))
}
async fn tpc_change_state_round_trip(
&mut self,
cx: &Cx,
operation: u32,
requested_state: u32,
xid: Option<&TpcXid<'_>>,
context: Option<&[u8]>,
) -> Result<TpcChangeStateResponse> {
cx.checkpoint()
.map_err(|err| Error::Runtime(err.to_string()))?;
let seq_num = next_ttc_sequence(&mut self.ttc_seq_num);
let payload = build_tpc_change_state_payload_with_seq(
seq_num,
operation,
requested_state,
0,
xid,
context,
);
send_data_packet_shared(cx, &self.write, &payload, self.sdu).await?;
let response = read_data_response(&mut self.read, cx, &self.write).await?;
self.note_parse(parse_tpc_change_state_response(
&response,
self.capabilities,
))
}
pub async fn tpc_begin(
&mut self,
cx: &Cx,
format_id: u32,
global_transaction_id: &[u8],
branch_qualifier: &[u8],
flags: u32,
timeout: u32,
) -> Result<()> {
let xid = TpcXid {
format_id,
global_transaction_id,
branch_qualifier,
};
let response = self
.tpc_switch_round_trip(cx, TNS_TPC_TXN_START, flags, timeout, Some(&xid), None)
.await?;
self.transaction_context = Some(response.context);
self.txn_in_progress = response.txn_in_progress;
Ok(())
}
pub async fn tpc_end(
&mut self,
cx: &Cx,
xid: Option<(u32, &[u8], &[u8])>,
flags: u32,
) -> Result<()> {
let xid = xid.map(|(format_id, gtid, bqual)| TpcXid {
format_id,
global_transaction_id: gtid,
branch_qualifier: bqual,
});
let context = self.transaction_context.clone();
let response = self
.tpc_switch_round_trip(
cx,
TNS_TPC_TXN_DETACH,
flags,
0,
xid.as_ref(),
context.as_deref(),
)
.await?;
self.txn_in_progress = response.txn_in_progress;
self.transaction_context = None;
Ok(())
}
pub async fn tpc_prepare(&mut self, cx: &Cx, xid: Option<(u32, &[u8], &[u8])>) -> Result<bool> {
let xid = xid.map(|(format_id, gtid, bqual)| TpcXid {
format_id,
global_transaction_id: gtid,
branch_qualifier: bqual,
});
let context = self.transaction_context.clone();
let response = self
.tpc_change_state_round_trip(
cx,
TNS_TPC_TXN_PREPARE,
TNS_TPC_TXN_STATE_PREPARE,
xid.as_ref(),
context.as_deref(),
)
.await?;
self.txn_in_progress = response.txn_in_progress;
match response.state {
TNS_TPC_TXN_STATE_REQUIRES_COMMIT => Ok(true),
TNS_TPC_TXN_STATE_READ_ONLY => Ok(false),
other => Err(Error::UnknownTransactionState(other)),
}
}
pub async fn tpc_commit(
&mut self,
cx: &Cx,
xid: Option<(u32, &[u8], &[u8])>,
one_phase: bool,
) -> Result<()> {
let xid = xid.map(|(format_id, gtid, bqual)| TpcXid {
format_id,
global_transaction_id: gtid,
branch_qualifier: bqual,
});
let requested_state = if one_phase {
TNS_TPC_TXN_STATE_READ_ONLY
} else {
TNS_TPC_TXN_STATE_COMMITTED
};
let context = self.transaction_context.clone();
let response = self
.tpc_change_state_round_trip(
cx,
TNS_TPC_TXN_COMMIT,
requested_state,
xid.as_ref(),
context.as_deref(),
)
.await?;
self.txn_in_progress = response.txn_in_progress;
let state = response.state;
let ok = if one_phase {
state == TNS_TPC_TXN_STATE_READ_ONLY || state == TNS_TPC_TXN_STATE_COMMITTED
} else {
state == TNS_TPC_TXN_STATE_FORGOTTEN
};
if !ok {
return Err(Error::UnknownTransactionState(state));
}
self.transaction_context = None;
Ok(())
}
pub async fn tpc_rollback(&mut self, cx: &Cx, xid: Option<(u32, &[u8], &[u8])>) -> Result<()> {
let xid = xid.map(|(format_id, gtid, bqual)| TpcXid {
format_id,
global_transaction_id: gtid,
branch_qualifier: bqual,
});
let context = self.transaction_context.clone();
let response = self
.tpc_change_state_round_trip(
cx,
TNS_TPC_TXN_ABORT,
TNS_TPC_TXN_STATE_ABORTED,
xid.as_ref(),
context.as_deref(),
)
.await?;
self.txn_in_progress = response.txn_in_progress;
if response.state != TNS_TPC_TXN_STATE_ABORTED {
return Err(Error::UnknownTransactionState(response.state));
}
Ok(())
}
pub fn transaction_in_progress(&self) -> bool {
self.txn_in_progress
}
pub fn prepare_sessionless_suspend_on_success(&mut self) -> Result<()> {
match &mut self.sessionless_data {
None => Err(Error::SessionlessTransaction(SessionlessError::Inactive)),
Some(data) if data.started_on_server => Err(Error::SessionlessTransaction(
SessionlessError::DifferingMethods,
)),
Some(data) => {
if data.piggyback_pending {
data.operation |= TNS_TPC_TXN_POST_DETACH;
} else {
data.operation = TNS_TPC_TXN_POST_DETACH;
data.flags = TPC_TXN_FLAGS_SESSIONLESS;
data.piggyback_pending = true;
}
Ok(())
}
}
}
fn take_sessionless_piggyback(&mut self) -> Option<Vec<u8>> {
let data = self.sessionless_data.as_mut()?;
if !data.piggyback_pending {
return None;
}
data.piggyback_pending = false;
let xid = if data.operation & TNS_TPC_TXN_START != 0 {
Some(data.transaction_id.clone())
} else {
None
};
let flags = data.flags | TPC_TXN_FLAGS_SESSIONLESS;
let operation = data.operation;
let timeout = data.timeout;
let seq_num = next_ttc_sequence(&mut self.ttc_seq_num);
Some(build_sessionless_piggyback(
seq_num,
0,
operation,
flags,
timeout,
xid.as_deref(),
))
}
fn apply_sessionless_state(&mut self, state: Option<SessionlessTxnState>) {
match state {
Some(SessionlessTxnState::Unset) => {
self.sessionless_data = None;
self.txn_in_progress = false;
}
Some(SessionlessTxnState::Set { started_on_server }) => {
self.txn_in_progress = true;
match self.sessionless_data.as_mut() {
Some(data) => {
data.started_on_server = started_on_server;
data.piggyback_pending = false;
}
None => {
self.sessionless_data = Some(SessionlessData {
transaction_id: Vec::new(),
timeout: 0,
operation: TNS_TPC_TXN_START,
flags: 0,
piggyback_pending: false,
started_on_server,
});
}
}
}
None => {}
}
}
pub async fn execute_query(
&mut self,
cx: &Cx,
sql: &str,
prefetch_rows: u32,
) -> Result<QueryResult> {
let _span = obs_span!(
"oracledb.execute",
db.statement = %crate::obs::sql_digest(sql),
db.bind_count = 0u64,
db.rows_fetched = tracing::field::Empty,
);
cx.checkpoint()
.map_err(|err| Error::Runtime(err.to_string()))?;
self.drain_pending_cancel(cx).await?;
let close_piggyback = self.take_close_cursors_piggyback();
let seq_num = next_ttc_sequence(&mut self.ttc_seq_num);
let mut payload =
build_execute_payload_with_seq(sql, prefetch_rows, seq_num, statement_is_query(sql))?;
if let Some(mut piggyback_bytes) = close_piggyback {
piggyback_bytes.extend_from_slice(&payload);
payload = piggyback_bytes;
}
trace_query_bytes("EXECUTE query payload", &payload);
send_data_packet_shared(cx, &self.write, &payload, self.sdu).await?;
let response = self.read_response_cancellable(cx).await?;
trace_query_bytes("EXECUTE query response", &response);
let parsed = parse_query_response(&response, self.capabilities);
let result = self.note_parse(parsed)?;
obs_record!(_span, db.rows_fetched = result.rows.len() as u64);
self.remember_cursor_columns(&result);
Ok(result)
}
pub async fn execute_query_collect(
&mut self,
cx: &Cx,
sql: &str,
prefetch_rows: u32,
) -> Result<QueryResult> {
let mut result = self.execute_query(cx, sql, prefetch_rows).await?;
if !columns_require_define(&result.columns) || result.cursor_id == 0 {
return Ok(result);
}
if !result.rows.is_empty() {
return Ok(result);
}
let cursor_id = result.cursor_id;
let columns = result.columns.clone();
let fetched = self
.define_and_fetch_rows_with_columns(cx, cursor_id, prefetch_rows.max(1), &columns, None)
.await?;
result.rows = fetched.rows;
result.more_rows = fetched.more_rows;
if !fetched.columns.is_empty() {
result.columns = fetched.columns;
}
if result.cursor_id == 0 {
result.cursor_id = cursor_id;
}
Ok(result)
}
pub async fn execute_query_with_timeout(
&mut self,
cx: &Cx,
sql: &str,
prefetch_rows: u32,
timeout_ms: Option<u32>,
) -> Result<QueryResult> {
self.execute_query_call_timeout(cx, sql, prefetch_rows, timeout_ms)
.await
}
pub async fn execute_query_with_binds(
&mut self,
cx: &Cx,
sql: &str,
prefetch_rows: u32,
binds: &[BindValue],
) -> Result<QueryResult> {
let bind_rows = if binds.is_empty() {
Vec::new()
} else {
vec![binds.to_vec()]
};
self.execute_query_with_bind_rows_and_options(
cx,
sql,
prefetch_rows,
&bind_rows,
ExecuteOptions::default(),
)
.await
}
pub async fn execute_query_with_binds_and_timeout(
&mut self,
cx: &Cx,
sql: &str,
prefetch_rows: u32,
binds: &[BindValue],
timeout_ms: Option<u32>,
) -> Result<QueryResult> {
self.execute_query_with_binds_call_timeout(cx, sql, prefetch_rows, binds, timeout_ms)
.await
}
pub async fn query(
&mut self,
cx: &Cx,
sql: &str,
params: impl crate::IntoBinds,
) -> Result<QueryResult> {
let binds = params.into_binds();
self.execute_query_with_binds(cx, sql, 1, &binds).await
}
pub async fn query_named(
&mut self,
cx: &Cx,
sql: &str,
named_params: Vec<(String, BindValue)>,
) -> Result<QueryResult> {
let binds = crate::sql_convert::order_named_binds(sql, named_params);
self.execute_query_with_binds(cx, sql, 1, &binds).await
}
pub async fn execute_query_with_bind_rows(
&mut self,
cx: &Cx,
sql: &str,
prefetch_rows: u32,
bind_rows: &[Vec<BindValue>],
) -> Result<QueryResult> {
self.execute_query_with_bind_rows_and_options(
cx,
sql,
prefetch_rows,
bind_rows,
ExecuteOptions::default(),
)
.await
}
pub async fn execute_query_with_bind_rows_and_options(
&mut self,
cx: &Cx,
sql: &str,
prefetch_rows: u32,
bind_rows: &[Vec<BindValue>],
exec_options: ExecuteOptions,
) -> Result<QueryResult> {
match self
.execute_query_with_bind_rows_options_adjusted(
cx,
sql,
prefetch_rows,
bind_rows,
exec_options,
)
.await
{
Err(err) if refetch_retry_applies(&err) && statement_is_query(sql) => {
self.forget_fetch_metadata(sql);
self.execute_query_with_bind_rows_options_adjusted(
cx,
sql,
prefetch_rows,
bind_rows,
exec_options,
)
.await
}
other => other,
}
}
async fn execute_query_with_bind_rows_options_adjusted(
&mut self,
cx: &Cx,
sql: &str,
prefetch_rows: u32,
bind_rows: &[Vec<BindValue>],
exec_options: ExecuteOptions,
) -> Result<QueryResult> {
let _span = obs_span!(
"oracledb.execute",
db.statement = %crate::obs::sql_digest(sql),
db.bind_count = bind_rows.first().map_or(0, Vec::len) as u64,
db.bind_rows = bind_rows.len() as u64,
db.rows_fetched = tracing::field::Empty,
);
cx.checkpoint()
.map_err(|err| Error::Runtime(err.to_string()))?;
self.drain_pending_cancel(cx).await?;
let mut exec_options = exec_options;
if exec_options.suspend_on_success {
self.prepare_sessionless_suspend_on_success()?;
}
let use_cache = exec_options.cache_statement && !exec_options.parse_only;
let mut is_copy = false;
if exec_options.cursor_id == 0 && !exec_options.parse_only {
if use_cache {
if self.statement_is_in_use(sql) {
is_copy = true;
} else if let Some(cursor_id) = self.statement_cache_get(sql) {
exec_options.cursor_id = cursor_id;
}
} else if let Some(cursor_id) = self.statement_cache_take(sql) {
exec_options.cursor_id = cursor_id;
}
}
if exec_options.cursor_id != 0 && statement_is_query(sql) {
if let Some(columns) = self.cursor_columns.get(&exec_options.cursor_id) {
if columns.iter().any(|column| {
column.ora_type_num == oracledb_protocol::thin::ORA_TYPE_NUM_VECTOR
}) {
exec_options.no_prefetch = true;
}
}
}
let piggyback = self.take_close_cursors_piggyback();
if piggyback.is_none() {
let has_ref_cursor_output = bind_rows.iter().any(|row| {
row.iter().any(|value| {
matches!(
value,
BindValue::Output {
ora_type_num: oracledb_protocol::thin::ORA_TYPE_NUM_CURSOR,
..
}
)
})
});
if has_ref_cursor_output {
let _ = next_ttc_sequence(&mut self.ttc_seq_num);
}
}
let sessionless_piggyback = self.take_sessionless_piggyback();
let seq_num = next_ttc_sequence(&mut self.ttc_seq_num);
let mut payload = build_execute_payload_with_bind_rows_and_options_with_seq(
sql,
prefetch_rows,
seq_num,
statement_is_query(sql),
bind_rows,
exec_options,
)?;
if let Some(piggyback_bytes) = sessionless_piggyback {
let mut combined = piggyback_bytes;
combined.extend_from_slice(&payload);
payload = combined;
}
if let Some(mut piggyback_bytes) = piggyback {
piggyback_bytes.extend_from_slice(&payload);
payload = piggyback_bytes;
}
trace_query_bytes("EXECUTE query payload", &payload);
send_data_packet_shared(cx, &self.write, &payload, self.sdu).await?;
let response = self.read_flushing_out_binds_cancellable(cx).await?;
trace_query_bytes("EXECUTE query response", &response);
let known_columns = if exec_options.cursor_id != 0 {
self.cursor_columns
.get(&exec_options.cursor_id)
.cloned()
.unwrap_or_default()
} else {
Vec::new()
};
let parsed = parse_query_response_with_binds_options_and_columns(
&response,
self.capabilities,
bind_rows.first().map(Vec::as_slice).unwrap_or(&[]),
exec_options,
&known_columns,
);
match self.note_parse(parsed) {
Ok(result) => {
self.apply_sessionless_state(result.sessionless_txn_state);
if let Some(txn_in_progress) = result.txn_in_progress {
self.txn_in_progress = txn_in_progress;
}
if is_copy {
if result.cursor_id != 0 {
self.copied_cursors.insert(result.cursor_id);
}
} else if use_cache {
self.statement_cache_put(sql, result.cursor_id);
}
if result.cursor_id != 0 && statement_is_query(sql) && !exec_options.parse_only {
self.in_use_cursors.insert(result.cursor_id);
}
self.invalidate_bound_ref_cursors(bind_rows);
self.remember_cursor_columns(&result);
obs_record!(_span, db.rows_fetched = result.rows.len() as u64);
if exec_options.parse_only {
return Ok(result);
}
self.apply_refetch_metadata(cx, sql, result, prefetch_rows.max(2))
.await
}
Err(err) => {
if use_cache {
self.statement_cache_invalidate(sql, exec_options.cursor_id);
}
Err(err)
}
}
}
pub async fn execute_query_with_bind_rows_and_timeout(
&mut self,
cx: &Cx,
sql: &str,
prefetch_rows: u32,
bind_rows: &[Vec<BindValue>],
timeout_ms: Option<u32>,
) -> Result<QueryResult> {
self.execute_query_with_bind_rows_call_timeout(
cx,
sql,
prefetch_rows,
bind_rows,
timeout_ms,
)
.await
}
pub async fn execute_query_with_bind_rows_options_and_timeout(
&mut self,
cx: &Cx,
sql: &str,
prefetch_rows: u32,
bind_rows: &[Vec<BindValue>],
exec_options: ExecuteOptions,
timeout_ms: Option<u32>,
) -> Result<QueryResult> {
let Some(timeout_ms) = timeout_ms.filter(|value| *value > 0) else {
return self
.execute_query_with_bind_rows_and_options(
cx,
sql,
prefetch_rows,
bind_rows,
exec_options,
)
.await;
};
match time::timeout(
time::wall_now(),
Duration::from_millis(u64::from(timeout_ms)),
self.execute_query_with_bind_rows_and_options(
cx,
sql,
prefetch_rows,
bind_rows,
exec_options,
),
)
.await
{
Ok(result) => result,
Err(_) => self.recover_from_call_timeout(cx, timeout_ms).await,
}
}
async fn drain_pending_cancel(&mut self, cx: &Cx) -> Result<()> {
if !self.cancel_drain_pending.swap(false, Ordering::SeqCst) {
return Ok(());
}
match cancel_and_drain_wire(
&mut self.read,
cx,
&self.write,
BREAK_DRAIN_RECOVERY_TIMEOUT,
)
.await
{
Ok(()) => Ok(()),
Err(err) => {
self.dead = true;
Err(err)
}
}
}
async fn read_response_cancellable(&mut self, cx: &Cx) -> Result<Vec<u8>> {
let pending = Arc::clone(&self.cancel_drain_pending);
let mut guard = CancelDrainGuard::arm(&pending);
let response = read_data_response(&mut self.read, cx, &self.write).await?;
guard.disarm();
Ok(response)
}
async fn read_flushing_out_binds_cancellable(&mut self, cx: &Cx) -> Result<Vec<u8>> {
let pending = Arc::clone(&self.cancel_drain_pending);
let mut guard = CancelDrainGuard::arm(&pending);
let response =
read_data_response_flushing_out_binds(&mut self.read, cx, &self.write, self.sdu)
.await?;
guard.disarm();
Ok(response)
}
pub async fn fetch_rows(
&mut self,
cx: &Cx,
cursor_id: u32,
arraysize: u32,
previous_row: Option<&[Option<oracledb_protocol::thin::QueryValue>]>,
) -> Result<QueryResult> {
self.fetch_rows_with_columns(cx, cursor_id, arraysize, &[], previous_row)
.await
}
pub async fn fetch_rows_with_columns(
&mut self,
cx: &Cx,
cursor_id: u32,
arraysize: u32,
known_columns: &[ColumnMetadata],
previous_row: Option<&[Option<oracledb_protocol::thin::QueryValue>]>,
) -> Result<QueryResult> {
let _span = obs_span!(
"oracledb.fetch",
db.cursor_id = cursor_id as u64,
db.arraysize = arraysize as u64,
db.rows_fetched = tracing::field::Empty,
);
cx.checkpoint()
.map_err(|err| Error::Runtime(err.to_string()))?;
self.drain_pending_cancel(cx).await?;
let seq_num = next_ttc_sequence(&mut self.ttc_seq_num);
let payload = build_fetch_payload_with_seq(cursor_id, arraysize, seq_num);
trace_query_bytes("FETCH payload", &payload);
send_data_packet_shared(cx, &self.write, &payload, self.sdu).await?;
let profile = fetch_profile::enabled();
let read_start = profile.then(time::wall_now);
let response = self.read_response_cancellable(cx).await?;
if let Some(start) = read_start {
fetch_profile::add_read(time::wall_now().duration_since(start));
}
trace_query_bytes("FETCH response", &response);
let columns = self
.cursor_columns
.get(&cursor_id)
.cloned()
.unwrap_or_else(|| known_columns.to_vec());
let decode_start = profile.then(time::wall_now);
let parsed =
parse_fetch_response_with_context(&response, self.capabilities, &columns, previous_row);
if let Some(start) = decode_start {
fetch_profile::add_decode(time::wall_now().duration_since(start));
}
let result = self.note_parse(parsed)?;
obs_record!(_span, db.rows_fetched = result.rows.len() as u64);
self.remember_cursor_columns(&result);
Ok(result)
}
pub async fn fetch_rows_ref(
&mut self,
cx: &Cx,
cursor_id: u32,
arraysize: u32,
previous_row: Option<&[Option<oracledb_protocol::thin::QueryValue>]>,
) -> Result<BorrowedFetchResult> {
cx.checkpoint()
.map_err(|err| Error::Runtime(err.to_string()))?;
let seq_num = next_ttc_sequence(&mut self.ttc_seq_num);
let payload = build_fetch_payload_with_seq(cursor_id, arraysize, seq_num);
trace_query_bytes("FETCH payload", &payload);
send_data_packet_shared(cx, &self.write, &payload, self.sdu).await?;
let profile = fetch_profile::enabled();
let read_start = profile.then(time::wall_now);
let response = read_data_response(&mut self.read, cx, &self.write).await?;
if let Some(start) = read_start {
fetch_profile::add_read(time::wall_now().duration_since(start));
}
trace_query_bytes("FETCH response", &response);
let columns = self
.cursor_columns
.get(&cursor_id)
.cloned()
.unwrap_or_default();
let decode_start = profile.then(time::wall_now);
let parsed =
parse_query_response_borrowed(&response, self.capabilities, &columns, previous_row);
if let Some(start) = decode_start {
fetch_profile::add_decode(time::wall_now().duration_since(start));
}
let result = self.note_parse(parsed)?;
if cursor_id != 0 && !result.batch.columns().is_empty() {
self.cursor_columns
.insert(cursor_id, result.batch.columns().to_vec());
}
Ok(result)
}
pub async fn fetch_rows_request(
&mut self,
cx: &Cx,
cursor_id: u32,
arraysize: u32,
) -> Result<()> {
cx.checkpoint()
.map_err(|err| Error::Runtime(err.to_string()))?;
self.drain_pending_cancel(cx).await?;
let seq_num = next_ttc_sequence(&mut self.ttc_seq_num);
let payload = build_fetch_payload_with_seq(cursor_id, arraysize, seq_num);
trace_query_bytes("FETCH payload (prefetch)", &payload);
send_data_packet_shared(cx, &self.write, &payload, self.sdu).await?;
self.cancel_drain_pending.store(true, Ordering::SeqCst);
Ok(())
}
pub async fn fetch_rows_ref_response(
&mut self,
cx: &Cx,
cursor_id: u32,
previous_row: Option<&[Option<oracledb_protocol::thin::QueryValue>]>,
) -> Result<BorrowedFetchResult> {
cx.checkpoint()
.map_err(|err| Error::Runtime(err.to_string()))?;
let profile = fetch_profile::enabled();
let read_start = profile.then(time::wall_now);
let response = self.read_response_cancellable(cx).await?;
if let Some(start) = read_start {
fetch_profile::add_read(time::wall_now().duration_since(start));
}
self.cancel_drain_pending.store(false, Ordering::SeqCst);
trace_query_bytes("FETCH response (prefetch)", &response);
let columns = self
.cursor_columns
.get(&cursor_id)
.cloned()
.unwrap_or_default();
let decode_start = profile.then(time::wall_now);
let parsed =
parse_query_response_borrowed(&response, self.capabilities, &columns, previous_row);
if let Some(start) = decode_start {
fetch_profile::add_decode(time::wall_now().duration_since(start));
}
let result = self.note_parse(parsed)?;
if cursor_id != 0 && !result.batch.columns().is_empty() {
self.cursor_columns
.insert(cursor_id, result.batch.columns().to_vec());
}
Ok(result)
}
pub async fn for_each_row_ref<F>(
&mut self,
cx: &Cx,
sql: &str,
arraysize: u32,
mut callback: F,
) -> Result<()>
where
F: FnMut(&[Option<QueryValueRef<'_>>]) -> Result<()>,
{
let first = self
.execute_query_with_bind_rows(cx, sql, arraysize, &[])
.await?;
let cursor_id = first.cursor_id;
for row in &first.rows {
let refs: Vec<Option<QueryValueRef<'_>>> = row
.iter()
.map(|cell| cell.as_ref().map(QueryValueRef::Owned))
.collect();
callback(&refs)?;
}
let mut more_rows = first.more_rows;
let mut previous_row: Option<Vec<Option<oracledb_protocol::thin::QueryValue>>> =
first.rows.last().cloned();
if more_rows && cursor_id != 0 {
self.fetch_rows_request(cx, cursor_id, arraysize).await?;
}
while more_rows && cursor_id != 0 {
let result = self
.fetch_rows_ref_response(cx, cursor_id, previous_row.as_deref())
.await?;
let next_more = result.more_rows;
if next_more {
self.fetch_rows_request(cx, cursor_id, arraysize).await?;
}
let mut last_owned: Option<Vec<Option<oracledb_protocol::thin::QueryValue>>> = None;
result.batch.for_each_row_ref(|row| {
last_owned = Some(
row.iter()
.map(|cell| cell.map(|v| v.to_owned_value()))
.collect(),
);
callback(row)
})?;
if let Some(last) = last_owned {
previous_row = Some(last);
}
more_rows = next_more;
}
self.release_cursor(cursor_id);
Ok(())
}
pub async fn define_and_fetch_rows_with_columns(
&mut self,
cx: &Cx,
cursor_id: u32,
arraysize: u32,
define_columns: &[ColumnMetadata],
previous_row: Option<&[Option<oracledb_protocol::thin::QueryValue>]>,
) -> Result<QueryResult> {
cx.checkpoint()
.map_err(|err| Error::Runtime(err.to_string()))?;
let seq_num = next_ttc_sequence(&mut self.ttc_seq_num);
let payload =
build_define_fetch_payload_with_seq(cursor_id, arraysize, seq_num, define_columns)?;
trace_query_bytes("DEFINE FETCH payload", &payload);
send_data_packet_shared(cx, &self.write, &payload, self.sdu).await?;
let response = read_data_response(&mut self.read, cx, &self.write).await?;
trace_query_bytes("DEFINE FETCH response", &response);
let result = parse_fetch_response_with_context(
&response,
self.capabilities,
define_columns,
previous_row,
)
.map_err(Error::from)?;
self.cursor_columns
.insert(cursor_id, define_columns.to_vec());
self.remember_cursor_columns(&result);
Ok(result)
}
pub async fn scroll_cursor(
&mut self,
cx: &Cx,
sql: &str,
cursor_id: u32,
arraysize: u32,
fetch_orientation: u32,
fetch_pos: u32,
) -> Result<QueryResult> {
cx.checkpoint()
.map_err(|err| Error::Runtime(err.to_string()))?;
let exec_options = ExecuteOptions {
cursor_id,
scrollable: true,
scroll_operation: true,
fetch_orientation,
fetch_pos,
cache_statement: false,
..ExecuteOptions::default()
};
let piggyback = self.take_close_cursors_piggyback();
let seq_num = next_ttc_sequence(&mut self.ttc_seq_num);
let mut payload = build_execute_payload_with_bind_rows_and_options_with_seq(
sql,
arraysize,
seq_num,
true,
&[],
exec_options,
)?;
if let Some(mut piggyback_bytes) = piggyback {
piggyback_bytes.extend_from_slice(&payload);
payload = piggyback_bytes;
}
trace_query_bytes("SCROLL payload", &payload);
send_data_packet_shared(cx, &self.write, &payload, self.sdu).await?;
let response =
read_data_response_flushing_out_binds(&mut self.read, cx, &self.write, self.sdu)
.await?;
trace_query_bytes("SCROLL response", &response);
let known_columns = self
.cursor_columns
.get(&cursor_id)
.cloned()
.unwrap_or_default();
let parsed = parse_query_response_with_binds_options_and_columns(
&response,
self.capabilities,
&[],
exec_options,
&known_columns,
);
let result = self.note_parse(parsed)?;
self.remember_cursor_columns(&result);
Ok(result)
}
pub async fn read_lob(
&mut self,
cx: &Cx,
locator: &[u8],
offset: u64,
amount: u64,
) -> Result<LobReadResult> {
let _span = obs_span!(
"oracledb.lob",
db.operation = "read",
db.lob_offset = offset,
db.lob_amount = amount,
);
cx.checkpoint()
.map_err(|err| Error::Runtime(err.to_string()))?;
let seq_num = next_ttc_sequence(&mut self.ttc_seq_num);
let payload = build_lob_read_payload_with_seq(
locator,
offset,
amount,
seq_num,
self.capabilities.ttc_field_version,
)?;
trace_query_bytes("LOB READ payload", &payload);
send_data_packet_shared(cx, &self.write, &payload, self.sdu).await?;
let response = read_data_response(&mut self.read, cx, &self.write).await?;
trace_query_bytes("LOB READ response", &response);
self.note_parse(parse_lob_read_response(
&response,
self.capabilities,
locator,
))
}
pub async fn read_lob_with_timeout(
&mut self,
cx: &Cx,
locator: &[u8],
offset: u64,
amount: u64,
timeout_ms: Option<u32>,
) -> Result<LobReadResult> {
self.read_lob_call_timeout(cx, locator, offset, amount, timeout_ms)
.await
}
pub async fn aq_enq_one(
&mut self,
cx: &Cx,
queue: &AqQueueDesc,
props: &AqMsgProps,
enq_options: &AqEnqOptions,
) -> Result<Option<Vec<u8>>> {
cx.checkpoint()
.map_err(|err| Error::Runtime(err.to_string()))?;
let seq_num = next_ttc_sequence(&mut self.ttc_seq_num);
let payload = build_aq_enq_payload(
queue,
props,
enq_options,
seq_num,
self.capabilities.ttc_field_version,
self.supports_oson_long_fnames(),
)?;
trace_query_bytes("AQ ENQ payload", &payload);
send_data_packet_shared(cx, &self.write, &payload, self.sdu).await?;
let response = read_data_response(&mut self.read, cx, &self.write).await?;
trace_query_bytes("AQ ENQ response", &response);
self.note_parse(parse_aq_enq_response(&response, self.capabilities))
}
pub async fn aq_deq_one(
&mut self,
cx: &Cx,
queue: &AqQueueDesc,
deq_options: &AqDeqOptions,
) -> Result<AqDeqResult> {
cx.checkpoint()
.map_err(|err| Error::Runtime(err.to_string()))?;
let seq_num = next_ttc_sequence(&mut self.ttc_seq_num);
let payload = build_aq_deq_payload(
queue,
deq_options,
seq_num,
self.capabilities.ttc_field_version,
)?;
trace_query_bytes("AQ DEQ payload", &payload);
send_data_packet_shared(cx, &self.write, &payload, self.sdu).await?;
let response = read_data_response(&mut self.read, cx, &self.write).await?;
trace_query_bytes("AQ DEQ response", &response);
self.note_parse(parse_aq_deq_response(
&response,
self.capabilities,
&queue.kind,
))
}
pub async fn aq_enq_many(
&mut self,
cx: &Cx,
queue: &AqQueueDesc,
props_list: &[AqMsgProps],
enq_options: &AqEnqOptions,
) -> Result<Vec<Vec<u8>>> {
cx.checkpoint()
.map_err(|err| Error::Runtime(err.to_string()))?;
let seq_num = next_ttc_sequence(&mut self.ttc_seq_num);
let payload = build_aq_array_enq_payload(
queue,
props_list,
enq_options,
seq_num,
self.capabilities.ttc_field_version,
self.supports_oson_long_fnames(),
)?;
trace_query_bytes("AQ ARRAY ENQ payload", &payload);
send_data_packet_shared(cx, &self.write, &payload, self.sdu).await?;
let response = read_data_response(&mut self.read, cx, &self.write).await?;
trace_query_bytes("AQ ARRAY ENQ response", &response);
let result: AqArrayResult = self.note_parse(parse_aq_array_response(
&response,
self.capabilities,
TNS_AQ_ARRAY_ENQ,
props_list.len() as u32,
&queue.kind,
))?;
Ok(result.enq_msgids)
}
pub async fn aq_deq_many(
&mut self,
cx: &Cx,
queue: &AqQueueDesc,
deq_options: &AqDeqOptions,
max_num_messages: u32,
) -> Result<Vec<oracledb_protocol::thin::aq::AqDeqMessage>> {
cx.checkpoint()
.map_err(|err| Error::Runtime(err.to_string()))?;
let seq_num = next_ttc_sequence(&mut self.ttc_seq_num);
let payload = build_aq_array_deq_payload(
queue,
deq_options,
max_num_messages,
seq_num,
self.capabilities.ttc_field_version,
)?;
trace_query_bytes("AQ ARRAY DEQ payload", &payload);
send_data_packet_shared(cx, &self.write, &payload, self.sdu).await?;
let response = read_data_response(&mut self.read, cx, &self.write).await?;
trace_query_bytes("AQ ARRAY DEQ response", &response);
let result: AqArrayResult = self.note_parse(parse_aq_array_response(
&response,
self.capabilities,
TNS_AQ_ARRAY_DEQ,
max_num_messages,
&queue.kind,
))?;
Ok(result.deq_messages)
}
pub async fn create_temp_lob(
&mut self,
cx: &Cx,
ora_type_num: u8,
csfrm: u8,
) -> Result<LobReadResult> {
cx.checkpoint()
.map_err(|err| Error::Runtime(err.to_string()))?;
let seq_num = next_ttc_sequence(&mut self.ttc_seq_num);
let payload = build_lob_create_temp_payload_with_seq(
ora_type_num,
csfrm,
seq_num,
self.capabilities.ttc_field_version,
)?;
trace_query_bytes("LOB CREATE TEMP payload", &payload);
send_data_packet_shared(cx, &self.write, &payload, self.sdu).await?;
let response = read_data_response(&mut self.read, cx, &self.write).await?;
trace_query_bytes("LOB CREATE TEMP response", &response);
self.note_parse(parse_lob_create_temp_response(&response, self.capabilities))
}
pub async fn write_lob(
&mut self,
cx: &Cx,
locator: &[u8],
offset: u64,
data: &[u8],
) -> Result<LobReadResult> {
let _span = obs_span!(
"oracledb.lob",
db.operation = "write",
db.lob_offset = offset,
db.lob_bytes = data.len() as u64,
);
cx.checkpoint()
.map_err(|err| Error::Runtime(err.to_string()))?;
let seq_num = next_ttc_sequence(&mut self.ttc_seq_num);
let payload = build_lob_write_payload_with_seq(
locator,
offset,
data,
seq_num,
self.capabilities.ttc_field_version,
)?;
trace_query_bytes("LOB WRITE payload", &payload);
send_data_packet_shared(cx, &self.write, &payload, self.sdu).await?;
let response = read_data_response(&mut self.read, cx, &self.write).await?;
trace_query_bytes("LOB WRITE response", &response);
self.note_parse(parse_lob_write_response(
&response,
self.capabilities,
locator,
))
}
pub async fn write_lob_with_timeout(
&mut self,
cx: &Cx,
locator: &[u8],
offset: u64,
data: &[u8],
timeout_ms: Option<u32>,
) -> Result<LobReadResult> {
self.write_lob_call_timeout(cx, locator, offset, data, timeout_ms)
.await
}
pub async fn trim_lob(
&mut self,
cx: &Cx,
locator: &[u8],
new_size: u64,
) -> Result<LobReadResult> {
cx.checkpoint()
.map_err(|err| Error::Runtime(err.to_string()))?;
let seq_num = next_ttc_sequence(&mut self.ttc_seq_num);
let payload = build_lob_trim_payload_with_seq(
locator,
new_size,
seq_num,
self.capabilities.ttc_field_version,
)?;
trace_query_bytes("LOB TRIM payload", &payload);
send_data_packet_shared(cx, &self.write, &payload, self.sdu).await?;
let response = read_data_response(&mut self.read, cx, &self.write).await?;
trace_query_bytes("LOB TRIM response", &response);
self.note_parse(parse_lob_trim_response(
&response,
self.capabilities,
locator,
))
}
pub async fn trim_lob_with_timeout(
&mut self,
cx: &Cx,
locator: &[u8],
new_size: u64,
timeout_ms: Option<u32>,
) -> Result<LobReadResult> {
self.trim_lob_call_timeout(cx, locator, new_size, timeout_ms)
.await
}
pub async fn free_temp_lobs(&mut self, cx: &Cx, locators: &[Vec<u8>]) -> Result<()> {
cx.checkpoint()
.map_err(|err| Error::Runtime(err.to_string()))?;
if locators.is_empty() {
return Ok(());
}
let returned_parameter_len = locators.iter().map(Vec::len).sum();
let seq_num = next_ttc_sequence(&mut self.ttc_seq_num);
let payload = build_lob_free_temp_payload_with_seq(
locators,
seq_num,
self.capabilities.ttc_field_version,
)?;
trace_query_bytes("LOB FREE TEMP payload", &payload);
send_data_packet_shared(cx, &self.write, &payload, self.sdu).await?;
let response = read_data_response(&mut self.read, cx, &self.write).await?;
trace_query_bytes("LOB FREE TEMP response", &response);
self.note_parse(parse_lob_free_temp_response(
&response,
self.capabilities,
returned_parameter_len,
))
}
pub async fn free_temp_lobs_with_timeout(
&mut self,
cx: &Cx,
locators: &[Vec<u8>],
timeout_ms: Option<u32>,
) -> Result<()> {
self.free_temp_lobs_call_timeout(cx, locators, timeout_ms)
.await
}
async fn execute_query_call_timeout(
&mut self,
cx: &Cx,
sql: &str,
prefetch_rows: u32,
timeout_ms: Option<u32>,
) -> Result<QueryResult> {
let Some(timeout_ms) = timeout_ms.filter(|value| *value > 0) else {
return self.execute_query(cx, sql, prefetch_rows).await;
};
match time::timeout(
time::wall_now(),
Duration::from_millis(u64::from(timeout_ms)),
self.execute_query(cx, sql, prefetch_rows),
)
.await
{
Ok(result) => result,
Err(_) => self.recover_from_call_timeout(cx, timeout_ms).await,
}
}
async fn execute_query_with_binds_call_timeout(
&mut self,
cx: &Cx,
sql: &str,
prefetch_rows: u32,
binds: &[BindValue],
timeout_ms: Option<u32>,
) -> Result<QueryResult> {
let Some(timeout_ms) = timeout_ms.filter(|value| *value > 0) else {
return self
.execute_query_with_binds(cx, sql, prefetch_rows, binds)
.await;
};
match time::timeout(
time::wall_now(),
Duration::from_millis(u64::from(timeout_ms)),
self.execute_query_with_binds(cx, sql, prefetch_rows, binds),
)
.await
{
Ok(result) => result,
Err(_) => self.recover_from_call_timeout(cx, timeout_ms).await,
}
}
async fn execute_query_with_bind_rows_call_timeout(
&mut self,
cx: &Cx,
sql: &str,
prefetch_rows: u32,
bind_rows: &[Vec<BindValue>],
timeout_ms: Option<u32>,
) -> Result<QueryResult> {
let Some(timeout_ms) = timeout_ms.filter(|value| *value > 0) else {
return self
.execute_query_with_bind_rows(cx, sql, prefetch_rows, bind_rows)
.await;
};
match time::timeout(
time::wall_now(),
Duration::from_millis(u64::from(timeout_ms)),
self.execute_query_with_bind_rows(cx, sql, prefetch_rows, bind_rows),
)
.await
{
Ok(result) => result,
Err(_) => self.recover_from_call_timeout(cx, timeout_ms).await,
}
}
async fn read_lob_call_timeout(
&mut self,
cx: &Cx,
locator: &[u8],
offset: u64,
amount: u64,
timeout_ms: Option<u32>,
) -> Result<LobReadResult> {
let Some(timeout_ms) = timeout_ms.filter(|value| *value > 0) else {
return self.read_lob(cx, locator, offset, amount).await;
};
match time::timeout(
time::wall_now(),
Duration::from_millis(u64::from(timeout_ms)),
self.read_lob(cx, locator, offset, amount),
)
.await
{
Ok(result) => result,
Err(_) => self.recover_from_call_timeout(cx, timeout_ms).await,
}
}
async fn write_lob_call_timeout(
&mut self,
cx: &Cx,
locator: &[u8],
offset: u64,
data: &[u8],
timeout_ms: Option<u32>,
) -> Result<LobReadResult> {
let Some(timeout_ms) = timeout_ms.filter(|value| *value > 0) else {
return self.write_lob(cx, locator, offset, data).await;
};
match time::timeout(
time::wall_now(),
Duration::from_millis(u64::from(timeout_ms)),
self.write_lob(cx, locator, offset, data),
)
.await
{
Ok(result) => result,
Err(_) => self.recover_from_call_timeout(cx, timeout_ms).await,
}
}
async fn trim_lob_call_timeout(
&mut self,
cx: &Cx,
locator: &[u8],
new_size: u64,
timeout_ms: Option<u32>,
) -> Result<LobReadResult> {
let Some(timeout_ms) = timeout_ms.filter(|value| *value > 0) else {
return self.trim_lob(cx, locator, new_size).await;
};
match time::timeout(
time::wall_now(),
Duration::from_millis(u64::from(timeout_ms)),
self.trim_lob(cx, locator, new_size),
)
.await
{
Ok(result) => result,
Err(_) => self.recover_from_call_timeout(cx, timeout_ms).await,
}
}
async fn free_temp_lobs_call_timeout(
&mut self,
cx: &Cx,
locators: &[Vec<u8>],
timeout_ms: Option<u32>,
) -> Result<()> {
let Some(timeout_ms) = timeout_ms.filter(|value| *value > 0) else {
return self.free_temp_lobs(cx, locators).await;
};
match time::timeout(
time::wall_now(),
Duration::from_millis(u64::from(timeout_ms)),
self.free_temp_lobs(cx, locators),
)
.await
{
Ok(result) => result,
Err(_) => self.recover_from_call_timeout(cx, timeout_ms).await,
}
}
pub async fn direct_path_prepare(
&mut self,
cx: &Cx,
schema_name: &str,
table_name: &str,
column_names: &[String],
) -> Result<oracledb_protocol::dpl::DirectPathPrepareResult> {
cx.checkpoint()
.map_err(|err| Error::Runtime(err.to_string()))?;
let seq_num = next_ttc_sequence(&mut self.ttc_seq_num);
let payload = oracledb_protocol::dpl::build_direct_path_prepare_payload(
schema_name,
table_name,
column_names,
seq_num,
)?;
trace_query_bytes("DIRECT PATH PREPARE payload", &payload);
send_data_packet_shared(cx, &self.write, &payload, self.sdu).await?;
let response = read_data_response(&mut self.read, cx, &self.write).await?;
trace_query_bytes("DIRECT PATH PREPARE response", &response);
oracledb_protocol::dpl::parse_direct_path_prepare_response(&response, self.capabilities)
.map_err(Error::from)
}
pub async fn direct_path_load_stream(
&mut self,
cx: &Cx,
cursor_id: u16,
stream: &oracledb_protocol::dpl::DirectPathStream,
) -> Result<()> {
cx.checkpoint()
.map_err(|err| Error::Runtime(err.to_string()))?;
let seq_num = next_ttc_sequence(&mut self.ttc_seq_num);
let payload = oracledb_protocol::dpl::build_direct_path_load_stream_payload(
cursor_id, stream, seq_num,
)?;
trace_query_bytes("DIRECT PATH LOAD STREAM payload", &payload);
send_data_packet_shared(cx, &self.write, &payload, self.sdu).await?;
let response = read_data_response(&mut self.read, cx, &self.write).await?;
trace_query_bytes("DIRECT PATH LOAD STREAM response", &response);
oracledb_protocol::dpl::parse_direct_path_simple_response(&response, self.capabilities)
.map_err(Error::from)
}
pub async fn direct_path_op(&mut self, cx: &Cx, cursor_id: u16, op_code: u32) -> Result<()> {
cx.checkpoint()
.map_err(|err| Error::Runtime(err.to_string()))?;
let seq_num = next_ttc_sequence(&mut self.ttc_seq_num);
let payload =
oracledb_protocol::dpl::build_direct_path_op_payload(cursor_id, op_code, seq_num);
trace_query_bytes("DIRECT PATH OP payload", &payload);
send_data_packet_shared(cx, &self.write, &payload, self.sdu).await?;
let response = read_data_response(&mut self.read, cx, &self.write).await?;
trace_query_bytes("DIRECT PATH OP response", &response);
oracledb_protocol::dpl::parse_direct_path_simple_response(&response, self.capabilities)
.map_err(Error::from)
}
pub async fn direct_path_load(
&mut self,
cx: &Cx,
schema_name: &str,
table_name: &str,
column_names: &[String],
rows: &[Vec<oracledb_protocol::dpl::DirectPathColumnValue>],
batch_size: u32,
) -> Result<()> {
let prepare = self
.direct_path_prepare(cx, schema_name, table_name, column_names)
.await?;
let load_result = self
.direct_path_load_batches(cx, &prepare, rows, batch_size)
.await;
let op_code = if load_result.is_ok() {
oracledb_protocol::dpl::TNS_DP_OP_FINISH
} else {
oracledb_protocol::dpl::TNS_DP_OP_ABORT
};
let op_result = self.direct_path_op(cx, prepare.cursor_id, op_code).await;
load_result?;
op_result
}
pub async fn direct_path_load_prepared(
&mut self,
cx: &Cx,
prepare: &oracledb_protocol::dpl::DirectPathPrepareResult,
rows: &[Vec<oracledb_protocol::dpl::DirectPathColumnValue>],
batch_size: u32,
) -> Result<()> {
let load_result = self
.direct_path_load_batches(cx, prepare, rows, batch_size)
.await;
let op_code = if load_result.is_ok() {
oracledb_protocol::dpl::TNS_DP_OP_FINISH
} else {
oracledb_protocol::dpl::TNS_DP_OP_ABORT
};
let op_result = self.direct_path_op(cx, prepare.cursor_id, op_code).await;
load_result?;
op_result
}
async fn direct_path_load_batches(
&mut self,
cx: &Cx,
prepare: &oracledb_protocol::dpl::DirectPathPrepareResult,
rows: &[Vec<oracledb_protocol::dpl::DirectPathColumnValue>],
batch_size: u32,
) -> Result<()> {
for row in rows {
if row.len() != prepare.column_metadata.len() {
return Err(oracledb_protocol::ProtocolError::TtcDecode(
"direct path row width does not match column metadata",
)
.into());
}
}
let mut state =
oracledb_protocol::dpl::BatchLoadState::for_rows(rows.len() as u64, batch_size)?;
let mut row_num: u64 = 1;
while !state.is_done() {
let start = usize::try_from(state.offset()).map_err(|_| {
oracledb_protocol::ProtocolError::TtcDecode("direct path offset overflow")
})?;
let end = start + state.num_rows() as usize;
let stream = oracledb_protocol::dpl::encode_direct_path_rows(
&prepare.column_metadata,
&rows[start..end],
row_num,
)?;
row_num += (end - start) as u64;
self.direct_path_load_stream(cx, prepare.cursor_id, &stream)
.await?;
state.next_batch();
}
Ok(())
}
async fn break_and_drain(&mut self, cx: &Cx) -> Result<()> {
match break_and_drain_wire(
&mut self.read,
cx,
&self.write,
BREAK_DRAIN_RECOVERY_TIMEOUT,
)
.await
{
Ok(()) => Ok(()),
Err(err) => {
self.dead = true;
Err(err)
}
}
}
async fn recover_from_call_timeout<T>(&mut self, cx: &Cx, timeout_ms: u32) -> Result<T> {
match self.break_and_drain(cx).await {
Ok(()) => Err(Error::CallTimeout(timeout_ms)),
Err(closed) => Err(closed),
}
}
async fn drain_cancel_response(&mut self, cx: &Cx) -> Result<()> {
match drain_cancel_wire(
&mut self.read,
cx,
&self.write,
BREAK_DRAIN_RECOVERY_TIMEOUT,
)
.await
{
Ok(()) => Ok(()),
Err(err) => {
self.dead = true;
Err(err)
}
}
}
pub async fn cancel(&mut self, cx: &Cx) -> Result<()> {
match cancel_and_drain_wire(
&mut self.read,
cx,
&self.write,
BREAK_DRAIN_RECOVERY_TIMEOUT,
)
.await
{
Ok(()) => Ok(()),
Err(err) => {
self.dead = true;
Err(err)
}
}
}
pub fn supports_oob(&self) -> bool {
self.supports_oob
}
fn remember_cursor_columns(&mut self, result: &QueryResult) {
if result.cursor_id != 0 && !result.columns.is_empty() {
if self.cursor_columns.get(&result.cursor_id) == Some(&result.columns) {
return;
}
self.cursor_columns
.insert(result.cursor_id, result.columns.clone());
}
}
fn remember_fetch_metadata(&mut self, sql: &str, columns: &[ColumnMetadata]) {
const FETCH_METADATA_RETENTION_CAP: usize = 100;
if !self.fetch_metadata_by_sql.contains_key(sql) {
if self.fetch_metadata_order.len() >= FETCH_METADATA_RETENTION_CAP {
if let Some(oldest) = self.fetch_metadata_order.pop_front() {
self.fetch_metadata_by_sql.remove(&oldest);
}
}
self.fetch_metadata_order.push_back(sql.to_string());
}
self.fetch_metadata_by_sql
.insert(sql.to_string(), columns.to_vec());
}
fn forget_fetch_metadata(&mut self, sql: &str) -> bool {
if self.fetch_metadata_by_sql.remove(sql).is_some() {
self.fetch_metadata_order.retain(|entry| entry != sql);
return true;
}
false
}
async fn apply_refetch_metadata(
&mut self,
cx: &Cx,
sql: &str,
mut result: QueryResult,
arraysize: u32,
) -> Result<QueryResult> {
if result.columns.is_empty() {
return Ok(result);
}
if let Some(previous_columns) = self.fetch_metadata_by_sql.get(sql) {
let mut adjusted = result.columns.clone();
let mut any_adjusted = false;
for (index, column) in adjusted.iter_mut().enumerate() {
if let Some(previous) = previous_columns.get(index) {
any_adjusted |= adjust_refetch_metadata(previous, column);
}
}
if any_adjusted && result.cursor_id != 0 {
let cursor_id = result.cursor_id;
let mut redefined = self
.define_and_fetch_rows_with_columns(
cx,
cursor_id,
arraysize.max(1),
&adjusted,
None,
)
.await?;
if redefined.columns.is_empty() {
redefined.columns = adjusted;
}
if redefined.cursor_id == 0 {
redefined.cursor_id = cursor_id;
}
result = redefined;
}
}
self.remember_fetch_metadata(sql, &result.columns);
Ok(result)
}
fn statement_cache_get(&mut self, sql: &str) -> Option<u32> {
let index = self
.statement_cache
.iter()
.position(|(cached_sql, _)| cached_sql == sql)?;
let cursor_id = self.statement_cache[index].1;
if cursor_id != 0 && self.in_use_cursors.contains(&cursor_id) {
return None;
}
let entry = self.statement_cache.remove(index);
self.statement_cache.push(entry);
Some(cursor_id)
}
fn statement_cache_take(&mut self, sql: &str) -> Option<u32> {
let index = self
.statement_cache
.iter()
.position(|(cached_sql, _)| cached_sql == sql)?;
Some(self.statement_cache.remove(index).1)
}
fn statement_cache_put(&mut self, sql: &str, cursor_id: u32) {
if cursor_id == 0 {
return;
}
if let Some(index) = self
.statement_cache
.iter()
.position(|(cached_sql, _)| cached_sql == sql)
{
let (_, cached_id) = self.statement_cache.remove(index);
if cached_id != 0 && cached_id != cursor_id {
self.cursors_to_close.push(cached_id);
}
}
self.statement_cache.push((sql.to_string(), cursor_id));
while self.statement_cache.len() > STATEMENT_CACHE_SIZE {
let (_, evicted_id) = self.statement_cache.remove(0);
if evicted_id != 0 {
self.cursors_to_close.push(evicted_id);
}
}
}
fn invalidate_bound_ref_cursors(&mut self, bind_rows: &[Vec<BindValue>]) {
for row in bind_rows {
for value in row {
if let BindValue::Cursor { cursor_id } = value {
if *cursor_id == 0 {
continue;
}
self.statement_cache
.retain(|(_, cached_id)| cached_id != cursor_id);
self.cursor_columns.remove(cursor_id);
}
}
}
}
pub fn release_cursor(&mut self, cursor_id: u32) {
if cursor_id == 0 {
return;
}
self.in_use_cursors.remove(&cursor_id);
if self.copied_cursors.remove(&cursor_id) {
self.cursors_to_close.push(cursor_id);
self.cursor_columns.remove(&cursor_id);
}
}
pub fn close_cursor(&mut self, cursor_id: u32) {
if cursor_id == 0 {
return;
}
self.in_use_cursors.remove(&cursor_id);
self.copied_cursors.remove(&cursor_id);
self.cursor_columns.remove(&cursor_id);
if !self.cursors_to_close.contains(&cursor_id) {
self.cursors_to_close.push(cursor_id);
}
}
fn statement_is_in_use(&self, sql: &str) -> bool {
self.statement_cache
.iter()
.find(|(cached_sql, _)| cached_sql == sql)
.is_some_and(|(_, cursor_id)| {
*cursor_id != 0 && self.in_use_cursors.contains(cursor_id)
})
}
fn statement_cache_invalidate(&mut self, sql: &str, cursor_id: u32) {
if let Some(index) = self
.statement_cache
.iter()
.position(|(cached_sql, _)| cached_sql == sql)
{
self.statement_cache.remove(index);
}
if cursor_id != 0 {
self.cursors_to_close.push(cursor_id);
self.cursor_columns.remove(&cursor_id);
self.in_use_cursors.remove(&cursor_id);
self.copied_cursors.remove(&cursor_id);
}
}
fn take_close_cursors_piggyback(&mut self) -> Option<Vec<u8>> {
if self.cursors_to_close.is_empty() {
return None;
}
let cursor_ids = std::mem::take(&mut self.cursors_to_close);
let seq_num = next_ttc_sequence(&mut self.ttc_seq_num);
Some(oracledb_protocol::thin::build_close_cursors_piggyback(
&cursor_ids,
seq_num,
))
}
pub async fn close(mut self, cx: &Cx) -> Result<()> {
cx.checkpoint()
.map_err(|err| Error::Runtime(err.to_string()))?;
match time::timeout(time::wall_now(), Duration::from_secs(5), self.rollback(cx)).await {
Ok(result) => result?,
Err(_) => {
let eof = encode_packet(
TNS_PACKET_TYPE_DATA,
0,
Some(oracledb_protocol::thin::TNS_DATA_FLAGS_EOF),
&[],
PacketLengthWidth::Large32,
)?;
let _ = write_all_shared(cx, &self.write, &eof).await;
let _ = shutdown_write_shared(cx, &self.write).await;
return Ok(());
}
}
let seq_num = next_ttc_sequence(&mut self.ttc_seq_num);
send_data_packet_shared(
cx,
&self.write,
&build_function_payload_with_seq(TNS_FUNC_LOGOFF, seq_num),
self.sdu,
)
.await?;
if let Ok(response) = time::timeout(
time::wall_now(),
Duration::from_secs(5),
read_data_response(&mut self.read, cx, &self.write),
)
.await
{
let _ = response?;
}
let eof = encode_packet(
TNS_PACKET_TYPE_DATA,
0,
Some(oracledb_protocol::thin::TNS_DATA_FLAGS_EOF),
&[],
PacketLengthWidth::Large32,
)?;
write_all_shared(cx, &self.write, &eof).await?;
let _ = shutdown_write_shared(cx, &self.write).await;
Ok(())
}
pub async fn run_pipeline(
&mut self,
cx: &Cx,
requests: &[PipelineRequest],
continue_on_error: bool,
) -> Result<Vec<Vec<u8>>> {
cx.checkpoint()
.map_err(|err| Error::Runtime(err.to_string()))?;
if requests.is_empty() {
return Ok(Vec::new());
}
let pipeline_mode = if continue_on_error {
TNS_PIPELINE_MODE_CONTINUE_ON_ERROR
} else {
TNS_PIPELINE_MODE_ABORT_ON_ERROR
};
for (index, request) in requests.iter().enumerate() {
let token_num = index as u64 + 1;
let mut payload = Vec::new();
let mut first_packet_flags = 0u16;
if index == 0 {
if let Some(close_piggyback) = self.take_close_cursors_piggyback() {
payload.extend_from_slice(&close_piggyback);
}
let piggyback_seq = next_ttc_sequence(&mut self.ttc_seq_num);
payload.extend_from_slice(&build_begin_pipeline_piggyback(
piggyback_seq,
token_num,
pipeline_mode,
));
first_packet_flags |= TNS_DATA_FLAGS_BEGIN_PIPELINE;
}
let seq_num = next_ttc_sequence(&mut self.ttc_seq_num);
match request {
PipelineRequest::Execute {
sql,
bind_rows,
prefetch_rows,
} => payload.extend_from_slice(
&build_execute_payload_with_bind_rows_with_seq_and_token(
sql,
*prefetch_rows,
seq_num,
statement_is_query(sql),
bind_rows,
token_num,
)?,
),
PipelineRequest::Commit => {
payload.extend_from_slice(&build_function_payload_with_seq_and_token(
TNS_FUNC_COMMIT,
seq_num,
token_num,
));
}
}
trace_query_bytes("PIPELINE op payload", &payload);
send_data_packet_shared_with_flags(
cx,
&self.write,
&payload,
self.sdu,
first_packet_flags,
TNS_DATA_FLAGS_END_OF_REQUEST,
)
.await?;
}
let seq_num = next_ttc_sequence(&mut self.ttc_seq_num);
let end_payload = build_end_pipeline_payload_with_seq(seq_num);
trace_query_bytes("PIPELINE end payload", &end_payload);
send_data_packet_shared(cx, &self.write, &end_payload, self.sdu).await?;
let mut responses = Vec::with_capacity(requests.len() + 1);
for _ in 0..=requests.len() {
let response =
read_data_response_boundary(&mut self.read, cx, &self.write, true).await?;
trace_query_bytes("PIPELINE response", &response.payload);
responses.push(response.payload);
}
Ok(responses)
}
pub async fn run_pipeline_decoded(
&mut self,
cx: &Cx,
requests: &[PipelineRequest],
continue_on_error: bool,
) -> Result<Vec<Result<QueryResult>>> {
let raw = self.run_pipeline(cx, requests, continue_on_error).await?;
let mut decoded = Vec::with_capacity(requests.len());
for (index, request) in requests.iter().enumerate() {
let payload = &raw[index];
let outcome = match request {
PipelineRequest::Commit => {
match parse_plain_function_response(payload, self.capabilities) {
Ok(txn_in_progress) => Ok(QueryResult {
txn_in_progress: Some(txn_in_progress),
..QueryResult::default()
}),
Err(err) => Err(Error::Protocol(err)),
}
}
PipelineRequest::Execute { sql, bind_rows, .. } => {
parse_query_response_with_binds_options_and_columns(
payload,
self.capabilities,
bind_rows.first().map(Vec::as_slice).unwrap_or(&[]),
ExecuteOptions::default(),
&[],
)
.map_err(Error::Protocol)
.inspect(|result| {
self.remember_cursor_columns(result);
if result.cursor_id != 0 && statement_is_query(sql) {
self.in_use_cursors.insert(result.cursor_id);
}
})
}
};
if let Ok(result) = &outcome {
if let Some(txn_in_progress) = result.txn_in_progress {
self.txn_in_progress = txn_in_progress;
}
}
decoded.push(outcome);
}
Ok(decoded)
}
async fn send_function(&mut self, cx: &Cx, function_code: u8) -> Result<()> {
cx.checkpoint()
.map_err(|err| Error::Runtime(err.to_string()))?;
let seq_num = next_ttc_sequence(&mut self.ttc_seq_num);
send_data_packet_shared(
cx,
&self.write,
&build_function_payload_with_seq(function_code, seq_num),
self.sdu,
)
.await?;
let response = read_data_response(&mut self.read, cx, &self.write).await?;
let txn_in_progress =
self.note_parse(parse_plain_function_response(&response, self.capabilities))?;
self.txn_in_progress = txn_in_progress;
Ok(())
}
pub async fn supplement_json_column_metadata(
&mut self,
cx: &Cx,
columns: &mut [ColumnMetadata],
timeout_ms: Option<u32>,
) -> Result<()> {
let candidates = json_lob_probe_candidates(columns);
if candidates.is_empty() {
return Ok(());
}
for (index, column_name) in candidates {
let result = self
.execute_query_with_binds_and_timeout(
cx,
"select 1 \
from all_json_columns \
where owner = sys_context('USERENV', 'CURRENT_SCHEMA') \
and column_name = :1",
1,
&[BindValue::Text(column_name)],
timeout_ms,
)
.await?;
if !result.rows.is_empty() {
columns[index].is_json = true;
}
}
Ok(())
}
}
impl CancelHandle {
pub fn cancel(&mut self) -> Result<()> {
let runtime = build_io_runtime()?;
let write = Arc::clone(&self.write);
runtime.block_on(async {
let cx = Cx::current()
.ok_or_else(|| Error::Runtime("asupersync did not install an ambient Cx".into()))?;
send_marker_shared(&cx, &write, TNS_MARKER_TYPE_BREAK).await
})
}
}
pub struct BlockingConnection;
impl BlockingConnection {
pub fn connect(options: ConnectOptions) -> Result<Connection> {
let runtime = build_io_runtime()?;
runtime.block_on(async {
let cx = Cx::current()
.ok_or_else(|| Error::Runtime("asupersync did not install an ambient Cx".into()))?;
Connection::connect(&cx, options).await
})
}
pub fn ping(connection: &mut Connection) -> Result<()> {
let runtime = build_io_runtime()?;
runtime.block_on(async {
let cx = Cx::current()
.ok_or_else(|| Error::Runtime("asupersync did not install an ambient Cx".into()))?;
connection.ping(&cx).await
})
}
pub fn ping_with_timeout(connection: &mut Connection, timeout_ms: u32) -> Result<()> {
let runtime = build_io_runtime()?;
runtime.block_on(async {
let cx = Cx::current()
.ok_or_else(|| Error::Runtime("asupersync did not install an ambient Cx".into()))?;
connection.ping_with_timeout(&cx, timeout_ms).await
})
}
pub fn change_password(
connection: &mut Connection,
old_password: &str,
new_password: &str,
) -> Result<()> {
let runtime = build_io_runtime()?;
runtime.block_on(async {
let cx = Cx::current()
.ok_or_else(|| Error::Runtime("asupersync did not install an ambient Cx".into()))?;
connection
.change_password(&cx, old_password, new_password)
.await
})
}
pub fn commit(connection: &mut Connection) -> Result<()> {
let runtime = build_io_runtime()?;
runtime.block_on(async {
let cx = Cx::current()
.ok_or_else(|| Error::Runtime("asupersync did not install an ambient Cx".into()))?;
connection.commit(&cx).await
})
}
#[allow(clippy::too_many_arguments)]
pub fn subscribe_register(
connection: &mut Connection,
namespace: u32,
name: Option<&str>,
public_qos: u32,
operations: u32,
timeout: u32,
grouping_class: u8,
grouping_value: u32,
grouping_type: u8,
) -> Result<SubscribeResult> {
let runtime = build_io_runtime()?;
runtime.block_on(async {
let cx = Cx::current()
.ok_or_else(|| Error::Runtime("asupersync did not install an ambient Cx".into()))?;
connection
.subscribe_register(
&cx,
namespace,
name,
public_qos,
operations,
timeout,
grouping_class,
grouping_value,
grouping_type,
)
.await
})
}
#[allow(clippy::too_many_arguments)]
pub fn subscribe_unregister(
connection: &mut Connection,
registration_id: u64,
client_id: &[u8],
namespace: u32,
name: Option<&str>,
public_qos: u32,
operations: u32,
timeout: u32,
grouping_class: u8,
grouping_value: u32,
grouping_type: u8,
) -> Result<()> {
let runtime = build_io_runtime()?;
runtime.block_on(async {
let cx = Cx::current()
.ok_or_else(|| Error::Runtime("asupersync did not install an ambient Cx".into()))?;
connection
.subscribe_unregister(
&cx,
registration_id,
client_id,
namespace,
name,
public_qos,
operations,
timeout,
grouping_class,
grouping_value,
grouping_type,
)
.await
})
}
pub fn execute_query_for_registration(
connection: &mut Connection,
sql: &str,
registration_id: u64,
) -> Result<Option<u64>> {
let runtime = build_io_runtime()?;
runtime.block_on(async {
let cx = Cx::current()
.ok_or_else(|| Error::Runtime("asupersync did not install an ambient Cx".into()))?;
connection
.execute_query_for_registration(&cx, sql, registration_id)
.await
})
}
pub fn rollback(connection: &mut Connection) -> Result<()> {
let runtime = build_io_runtime()?;
runtime.block_on(async {
let cx = Cx::current()
.ok_or_else(|| Error::Runtime("asupersync did not install an ambient Cx".into()))?;
connection.rollback(&cx).await
})
}
pub fn begin_sessionless_transaction(
connection: &mut Connection,
transaction_id: &[u8],
timeout: u32,
defer_round_trip: bool,
) -> Result<()> {
let runtime = build_io_runtime()?;
runtime.block_on(async {
let cx = Cx::current()
.ok_or_else(|| Error::Runtime("asupersync did not install an ambient Cx".into()))?;
connection
.begin_sessionless_transaction(&cx, transaction_id, timeout, defer_round_trip)
.await
})
}
pub fn resume_sessionless_transaction(
connection: &mut Connection,
transaction_id: &[u8],
timeout: u32,
defer_round_trip: bool,
) -> Result<()> {
let runtime = build_io_runtime()?;
runtime.block_on(async {
let cx = Cx::current()
.ok_or_else(|| Error::Runtime("asupersync did not install an ambient Cx".into()))?;
connection
.resume_sessionless_transaction(&cx, transaction_id, timeout, defer_round_trip)
.await
})
}
pub fn suspend_sessionless_transaction(connection: &mut Connection) -> Result<()> {
let runtime = build_io_runtime()?;
runtime.block_on(async {
let cx = Cx::current()
.ok_or_else(|| Error::Runtime("asupersync did not install an ambient Cx".into()))?;
connection.suspend_sessionless_transaction(&cx).await
})
}
#[allow(clippy::too_many_arguments)]
pub fn tpc_begin(
connection: &mut Connection,
format_id: u32,
global_transaction_id: &[u8],
branch_qualifier: &[u8],
flags: u32,
timeout: u32,
) -> Result<()> {
let runtime = build_io_runtime()?;
runtime.block_on(async {
let cx = Cx::current()
.ok_or_else(|| Error::Runtime("asupersync did not install an ambient Cx".into()))?;
connection
.tpc_begin(
&cx,
format_id,
global_transaction_id,
branch_qualifier,
flags,
timeout,
)
.await
})
}
pub fn tpc_end(
connection: &mut Connection,
xid: Option<(u32, &[u8], &[u8])>,
flags: u32,
) -> Result<()> {
let runtime = build_io_runtime()?;
runtime.block_on(async {
let cx = Cx::current()
.ok_or_else(|| Error::Runtime("asupersync did not install an ambient Cx".into()))?;
connection.tpc_end(&cx, xid, flags).await
})
}
pub fn tpc_prepare(
connection: &mut Connection,
xid: Option<(u32, &[u8], &[u8])>,
) -> Result<bool> {
let runtime = build_io_runtime()?;
runtime.block_on(async {
let cx = Cx::current()
.ok_or_else(|| Error::Runtime("asupersync did not install an ambient Cx".into()))?;
connection.tpc_prepare(&cx, xid).await
})
}
pub fn tpc_commit(
connection: &mut Connection,
xid: Option<(u32, &[u8], &[u8])>,
one_phase: bool,
) -> Result<()> {
let runtime = build_io_runtime()?;
runtime.block_on(async {
let cx = Cx::current()
.ok_or_else(|| Error::Runtime("asupersync did not install an ambient Cx".into()))?;
connection.tpc_commit(&cx, xid, one_phase).await
})
}
pub fn tpc_rollback(
connection: &mut Connection,
xid: Option<(u32, &[u8], &[u8])>,
) -> Result<()> {
let runtime = build_io_runtime()?;
runtime.block_on(async {
let cx = Cx::current()
.ok_or_else(|| Error::Runtime("asupersync did not install an ambient Cx".into()))?;
connection.tpc_rollback(&cx, xid).await
})
}
pub fn execute_query(
connection: &mut Connection,
sql: &str,
prefetch_rows: u32,
) -> Result<QueryResult> {
let runtime = build_io_runtime()?;
runtime.block_on(async {
let cx = Cx::current()
.ok_or_else(|| Error::Runtime("asupersync did not install an ambient Cx".into()))?;
connection.execute_query(&cx, sql, prefetch_rows).await
})
}
pub fn execute_query_collect(
connection: &mut Connection,
sql: &str,
prefetch_rows: u32,
) -> Result<QueryResult> {
let runtime = build_io_runtime()?;
runtime.block_on(async {
let cx = Cx::current()
.ok_or_else(|| Error::Runtime("asupersync did not install an ambient Cx".into()))?;
connection
.execute_query_collect(&cx, sql, prefetch_rows)
.await
})
}
pub fn execute_query_with_timeout(
connection: &mut Connection,
sql: &str,
prefetch_rows: u32,
timeout_ms: Option<u32>,
) -> Result<QueryResult> {
let runtime = build_io_runtime()?;
runtime.block_on(async {
let cx = Cx::current()
.ok_or_else(|| Error::Runtime("asupersync did not install an ambient Cx".into()))?;
connection
.execute_query_call_timeout(&cx, sql, prefetch_rows, timeout_ms)
.await
})
}
pub fn execute_query_with_binds(
connection: &mut Connection,
sql: &str,
prefetch_rows: u32,
binds: &[BindValue],
) -> Result<QueryResult> {
let runtime = build_io_runtime()?;
runtime.block_on(async {
let cx = Cx::current()
.ok_or_else(|| Error::Runtime("asupersync did not install an ambient Cx".into()))?;
connection
.execute_query_with_binds(&cx, sql, prefetch_rows, binds)
.await
})
}
pub fn execute_query_with_binds_and_timeout(
connection: &mut Connection,
sql: &str,
prefetch_rows: u32,
binds: &[BindValue],
timeout_ms: Option<u32>,
) -> Result<QueryResult> {
let runtime = build_io_runtime()?;
runtime.block_on(async {
let cx = Cx::current()
.ok_or_else(|| Error::Runtime("asupersync did not install an ambient Cx".into()))?;
connection
.execute_query_with_binds_call_timeout(&cx, sql, prefetch_rows, binds, timeout_ms)
.await
})
}
pub fn query(
connection: &mut Connection,
sql: &str,
params: impl crate::IntoBinds,
) -> Result<QueryResult> {
let runtime = build_io_runtime()?;
runtime.block_on(async {
let cx = Cx::current()
.ok_or_else(|| Error::Runtime("asupersync did not install an ambient Cx".into()))?;
connection.query(&cx, sql, params).await
})
}
pub fn query_named(
connection: &mut Connection,
sql: &str,
named_params: Vec<(String, BindValue)>,
) -> Result<QueryResult> {
let runtime = build_io_runtime()?;
runtime.block_on(async {
let cx = Cx::current()
.ok_or_else(|| Error::Runtime("asupersync did not install an ambient Cx".into()))?;
connection.query_named(&cx, sql, named_params).await
})
}
pub fn execute_query_with_bind_rows(
connection: &mut Connection,
sql: &str,
prefetch_rows: u32,
bind_rows: &[Vec<BindValue>],
) -> Result<QueryResult> {
let runtime = build_io_runtime()?;
runtime.block_on(async {
let cx = Cx::current()
.ok_or_else(|| Error::Runtime("asupersync did not install an ambient Cx".into()))?;
connection
.execute_query_with_bind_rows(&cx, sql, prefetch_rows, bind_rows)
.await
})
}
pub fn execute_query_with_bind_rows_and_timeout(
connection: &mut Connection,
sql: &str,
prefetch_rows: u32,
bind_rows: &[Vec<BindValue>],
timeout_ms: Option<u32>,
) -> Result<QueryResult> {
let runtime = build_io_runtime()?;
runtime.block_on(async {
let cx = Cx::current()
.ok_or_else(|| Error::Runtime("asupersync did not install an ambient Cx".into()))?;
connection
.execute_query_with_bind_rows_call_timeout(
&cx,
sql,
prefetch_rows,
bind_rows,
timeout_ms,
)
.await
})
}
pub fn execute_query_with_bind_rows_options_and_timeout(
connection: &mut Connection,
sql: &str,
prefetch_rows: u32,
bind_rows: &[Vec<BindValue>],
exec_options: ExecuteOptions,
timeout_ms: Option<u32>,
) -> Result<QueryResult> {
let runtime = build_io_runtime()?;
runtime.block_on(async {
let cx = Cx::current()
.ok_or_else(|| Error::Runtime("asupersync did not install an ambient Cx".into()))?;
connection
.execute_query_with_bind_rows_options_and_timeout(
&cx,
sql,
prefetch_rows,
bind_rows,
exec_options,
timeout_ms,
)
.await
})
}
pub fn fetch_rows(
connection: &mut Connection,
cursor_id: u32,
arraysize: u32,
previous_row: Option<&[Option<oracledb_protocol::thin::QueryValue>]>,
) -> Result<QueryResult> {
let runtime = build_io_runtime()?;
runtime.block_on(async {
let cx = Cx::current()
.ok_or_else(|| Error::Runtime("asupersync did not install an ambient Cx".into()))?;
connection
.fetch_rows(&cx, cursor_id, arraysize, previous_row)
.await
})
}
pub fn fetch_rows_with_columns(
connection: &mut Connection,
cursor_id: u32,
arraysize: u32,
known_columns: &[ColumnMetadata],
previous_row: Option<&[Option<oracledb_protocol::thin::QueryValue>]>,
) -> Result<QueryResult> {
let runtime = build_io_runtime()?;
runtime.block_on(async {
let cx = Cx::current()
.ok_or_else(|| Error::Runtime("asupersync did not install an ambient Cx".into()))?;
connection
.fetch_rows_with_columns(&cx, cursor_id, arraysize, known_columns, previous_row)
.await
})
}
pub fn define_and_fetch_rows_with_columns(
connection: &mut Connection,
cursor_id: u32,
arraysize: u32,
define_columns: &[ColumnMetadata],
previous_row: Option<&[Option<oracledb_protocol::thin::QueryValue>]>,
) -> Result<QueryResult> {
let runtime = build_io_runtime()?;
runtime.block_on(async {
let cx = Cx::current()
.ok_or_else(|| Error::Runtime("asupersync did not install an ambient Cx".into()))?;
connection
.define_and_fetch_rows_with_columns(
&cx,
cursor_id,
arraysize,
define_columns,
previous_row,
)
.await
})
}
pub fn scroll_cursor(
connection: &mut Connection,
sql: &str,
cursor_id: u32,
arraysize: u32,
fetch_orientation: u32,
fetch_pos: u32,
) -> Result<QueryResult> {
let runtime = build_io_runtime()?;
runtime.block_on(async {
let cx = Cx::current()
.ok_or_else(|| Error::Runtime("asupersync did not install an ambient Cx".into()))?;
connection
.scroll_cursor(&cx, sql, cursor_id, arraysize, fetch_orientation, fetch_pos)
.await
})
}
pub fn read_lob(
connection: &mut Connection,
locator: &[u8],
offset: u64,
amount: u64,
) -> Result<LobReadResult> {
let runtime = build_io_runtime()?;
runtime.block_on(async {
let cx = Cx::current()
.ok_or_else(|| Error::Runtime("asupersync did not install an ambient Cx".into()))?;
connection.read_lob(&cx, locator, offset, amount).await
})
}
pub fn read_lob_with_timeout(
connection: &mut Connection,
locator: &[u8],
offset: u64,
amount: u64,
timeout_ms: Option<u32>,
) -> Result<LobReadResult> {
let runtime = build_io_runtime()?;
runtime.block_on(async {
let cx = Cx::current()
.ok_or_else(|| Error::Runtime("asupersync did not install an ambient Cx".into()))?;
connection
.read_lob_call_timeout(&cx, locator, offset, amount, timeout_ms)
.await
})
}
pub fn aq_enq_one(
connection: &mut Connection,
queue: &AqQueueDesc,
props: &AqMsgProps,
enq_options: &AqEnqOptions,
) -> Result<Option<Vec<u8>>> {
let runtime = build_io_runtime()?;
runtime.block_on(async {
let cx = Cx::current()
.ok_or_else(|| Error::Runtime("asupersync did not install an ambient Cx".into()))?;
connection.aq_enq_one(&cx, queue, props, enq_options).await
})
}
pub fn aq_deq_one(
connection: &mut Connection,
queue: &AqQueueDesc,
deq_options: &AqDeqOptions,
) -> Result<AqDeqResult> {
let runtime = build_io_runtime()?;
runtime.block_on(async {
let cx = Cx::current()
.ok_or_else(|| Error::Runtime("asupersync did not install an ambient Cx".into()))?;
connection.aq_deq_one(&cx, queue, deq_options).await
})
}
pub fn aq_enq_many(
connection: &mut Connection,
queue: &AqQueueDesc,
props_list: &[AqMsgProps],
enq_options: &AqEnqOptions,
) -> Result<Vec<Vec<u8>>> {
let runtime = build_io_runtime()?;
runtime.block_on(async {
let cx = Cx::current()
.ok_or_else(|| Error::Runtime("asupersync did not install an ambient Cx".into()))?;
connection
.aq_enq_many(&cx, queue, props_list, enq_options)
.await
})
}
pub fn aq_deq_many(
connection: &mut Connection,
queue: &AqQueueDesc,
deq_options: &AqDeqOptions,
max_num_messages: u32,
) -> Result<Vec<oracledb_protocol::thin::aq::AqDeqMessage>> {
let runtime = build_io_runtime()?;
runtime.block_on(async {
let cx = Cx::current()
.ok_or_else(|| Error::Runtime("asupersync did not install an ambient Cx".into()))?;
connection
.aq_deq_many(&cx, queue, deq_options, max_num_messages)
.await
})
}
pub fn create_temp_lob(
connection: &mut Connection,
ora_type_num: u8,
csfrm: u8,
) -> Result<LobReadResult> {
let runtime = build_io_runtime()?;
runtime.block_on(async {
let cx = Cx::current()
.ok_or_else(|| Error::Runtime("asupersync did not install an ambient Cx".into()))?;
connection.create_temp_lob(&cx, ora_type_num, csfrm).await
})
}
pub fn write_lob(
connection: &mut Connection,
locator: &[u8],
offset: u64,
data: &[u8],
) -> Result<LobReadResult> {
let runtime = build_io_runtime()?;
runtime.block_on(async {
let cx = Cx::current()
.ok_or_else(|| Error::Runtime("asupersync did not install an ambient Cx".into()))?;
connection.write_lob(&cx, locator, offset, data).await
})
}
pub fn write_lob_with_timeout(
connection: &mut Connection,
locator: &[u8],
offset: u64,
data: &[u8],
timeout_ms: Option<u32>,
) -> Result<LobReadResult> {
let runtime = build_io_runtime()?;
runtime.block_on(async {
let cx = Cx::current()
.ok_or_else(|| Error::Runtime("asupersync did not install an ambient Cx".into()))?;
connection
.write_lob_call_timeout(&cx, locator, offset, data, timeout_ms)
.await
})
}
pub fn trim_lob_with_timeout(
connection: &mut Connection,
locator: &[u8],
new_size: u64,
timeout_ms: Option<u32>,
) -> Result<LobReadResult> {
let runtime = build_io_runtime()?;
runtime.block_on(async {
let cx = Cx::current()
.ok_or_else(|| Error::Runtime("asupersync did not install an ambient Cx".into()))?;
connection
.trim_lob_call_timeout(&cx, locator, new_size, timeout_ms)
.await
})
}
pub fn free_temp_lobs_with_timeout(
connection: &mut Connection,
locators: &[Vec<u8>],
timeout_ms: Option<u32>,
) -> Result<()> {
let runtime = build_io_runtime()?;
runtime.block_on(async {
let cx = Cx::current()
.ok_or_else(|| Error::Runtime("asupersync did not install an ambient Cx".into()))?;
connection
.free_temp_lobs_call_timeout(&cx, locators, timeout_ms)
.await
})
}
pub fn direct_path_load(
connection: &mut Connection,
schema_name: &str,
table_name: &str,
column_names: &[String],
rows: &[Vec<oracledb_protocol::dpl::DirectPathColumnValue>],
batch_size: u32,
) -> Result<()> {
let runtime = build_io_runtime()?;
runtime.block_on(async {
let cx = Cx::current()
.ok_or_else(|| Error::Runtime("asupersync did not install an ambient Cx".into()))?;
connection
.direct_path_load(&cx, schema_name, table_name, column_names, rows, batch_size)
.await
})
}
pub fn run_pipeline(
connection: &mut Connection,
requests: &[PipelineRequest],
continue_on_error: bool,
) -> Result<Vec<Vec<u8>>> {
let runtime = build_io_runtime()?;
runtime.block_on(async {
let cx = Cx::current()
.ok_or_else(|| Error::Runtime("asupersync did not install an ambient Cx".into()))?;
connection
.run_pipeline(&cx, requests, continue_on_error)
.await
})
}
pub fn run_pipeline_decoded(
connection: &mut Connection,
requests: &[PipelineRequest],
continue_on_error: bool,
) -> Result<Vec<Result<QueryResult>>> {
let runtime = build_io_runtime()?;
runtime.block_on(async {
let cx = Cx::current()
.ok_or_else(|| Error::Runtime("asupersync did not install an ambient Cx".into()))?;
connection
.run_pipeline_decoded(&cx, requests, continue_on_error)
.await
})
}
pub fn direct_path_prepare(
connection: &mut Connection,
schema_name: &str,
table_name: &str,
column_names: &[String],
) -> Result<oracledb_protocol::dpl::DirectPathPrepareResult> {
let runtime = build_io_runtime()?;
runtime.block_on(async {
let cx = Cx::current()
.ok_or_else(|| Error::Runtime("asupersync did not install an ambient Cx".into()))?;
connection
.direct_path_prepare(&cx, schema_name, table_name, column_names)
.await
})
}
pub fn direct_path_load_prepared(
connection: &mut Connection,
prepare: &oracledb_protocol::dpl::DirectPathPrepareResult,
rows: &[Vec<oracledb_protocol::dpl::DirectPathColumnValue>],
batch_size: u32,
) -> Result<()> {
let runtime = build_io_runtime()?;
runtime.block_on(async {
let cx = Cx::current()
.ok_or_else(|| Error::Runtime("asupersync did not install an ambient Cx".into()))?;
connection
.direct_path_load_prepared(&cx, prepare, rows, batch_size)
.await
})
}
pub fn drain_cancel_response(connection: &mut Connection) -> Result<()> {
let runtime = build_io_runtime()?;
runtime.block_on(async {
let cx = Cx::current()
.ok_or_else(|| Error::Runtime("asupersync did not install an ambient Cx".into()))?;
connection.drain_cancel_response(&cx).await
})
}
pub fn supplement_json_column_metadata(
connection: &mut Connection,
columns: &mut [ColumnMetadata],
timeout_ms: Option<u32>,
) -> Result<()> {
let runtime = build_io_runtime()?;
runtime.block_on(async {
let cx = Cx::current()
.ok_or_else(|| Error::Runtime("asupersync did not install an ambient Cx".into()))?;
connection
.supplement_json_column_metadata(&cx, columns, timeout_ms)
.await
})
}
pub fn close(connection: Connection) -> Result<()> {
let runtime = build_io_runtime()?;
runtime.block_on(async {
let cx = Cx::current()
.ok_or_else(|| Error::Runtime("asupersync did not install an ambient Cx".into()))?;
connection.close(&cx).await
})
}
}
fn new_io_runtime() -> Result<Runtime> {
let reactor = reactor::create_reactor()?;
RuntimeBuilder::current_thread()
.with_reactor(reactor)
.build()
.map_err(|err| Error::Runtime(err.to_string()))
}
thread_local! {
static IO_RUNTIME: std::cell::RefCell<Option<Runtime>> =
const { std::cell::RefCell::new(None) };
}
fn build_io_runtime() -> Result<Runtime> {
IO_RUNTIME.with(|slot| {
if let Some(runtime) = slot.borrow().as_ref() {
return Ok(runtime.clone());
}
let runtime = new_io_runtime()?;
*slot.borrow_mut() = Some(runtime.clone());
Ok(runtime)
})
}
#[cfg(feature = "arrow")]
pub(crate) fn block_on_connection<F, Fut, T>(operation: F) -> Result<T>
where
F: FnOnce(Cx) -> Fut,
Fut: std::future::Future<Output = Result<T>>,
{
let runtime = build_io_runtime()?;
runtime.block_on(async {
let cx = Cx::current()
.ok_or_else(|| Error::Runtime("asupersync did not install an ambient Cx".into()))?;
operation(cx).await
})
}
#[derive(Clone, Debug, Eq, PartialEq)]
struct IncomingPacket {
packet_type: u8,
payload: Vec<u8>,
}
async fn lock_write<'a>(
cx: &Cx,
write: &'a SharedWriteHalf,
) -> Result<asupersync::sync::MutexGuard<'a, OracleWriteHalf>> {
write
.lock(cx)
.await
.map_err(|err| Error::Runtime(err.to_string()))
}
async fn write_all_shared(cx: &Cx, write: &SharedWriteHalf, packet: &[u8]) -> Result<()> {
let mut guard = lock_write(cx, write).await?;
guard.write_all(packet).await?;
guard.flush().await?;
Ok(())
}
async fn shutdown_write_shared(cx: &Cx, write: &SharedWriteHalf) -> Result<()> {
let mut guard = lock_write(cx, write).await?;
guard.shutdown().await?;
Ok(())
}
async fn send_data_packet_shared(
cx: &Cx,
write: &SharedWriteHalf,
payload: &[u8],
sdu: usize,
) -> Result<()> {
let mut guard = lock_write(cx, write).await?;
send_data_packet(&mut *guard, payload, sdu).await
}
async fn send_data_packet_shared_with_flags(
cx: &Cx,
write: &SharedWriteHalf,
payload: &[u8],
sdu: usize,
first_packet_flags: u16,
last_packet_flags: u16,
) -> Result<()> {
let mut guard = lock_write(cx, write).await?;
send_data_packet_with_flags(
&mut *guard,
payload,
sdu,
first_packet_flags,
last_packet_flags,
)
.await
}
async fn send_marker_shared(cx: &Cx, write: &SharedWriteHalf, marker_type: u8) -> Result<()> {
let mut guard = lock_write(cx, write).await?;
send_marker(&mut *guard, marker_type).await
}
async fn send_data_packet<W>(stream: &mut W, payload: &[u8], sdu: usize) -> Result<()>
where
W: AsyncWrite + Unpin,
{
send_data_packet_with_flags(stream, payload, sdu, 0, 0).await
}
async fn send_data_packet_with_flags<W>(
stream: &mut W,
payload: &[u8],
sdu: usize,
first_packet_flags: u16,
last_packet_flags: u16,
) -> Result<()>
where
W: AsyncWrite + Unpin,
{
let max_payload = sdu.saturating_sub(TNS_DATA_PACKET_OVERHEAD).max(1);
let chunk_count = payload.chunks(max_payload).len();
for (index, chunk) in payload.chunks(max_payload).enumerate() {
let mut flags = 0u16;
if index == 0 {
flags |= first_packet_flags;
}
if index + 1 == chunk_count {
flags |= last_packet_flags;
}
let packet = encode_packet(
TNS_PACKET_TYPE_DATA,
0,
Some(flags),
chunk,
PacketLengthWidth::Large32,
)?;
stream.write_all(&packet).await?;
}
stream.flush().await?;
Ok(())
}
struct DataResponse {
payload: Vec<u8>,
flush_out_binds: bool,
}
async fn read_data_response(
read: &mut OracleReadHalf,
cx: &Cx,
write: &SharedWriteHalf,
) -> Result<Vec<u8>> {
Ok(read_data_response_boundary(read, cx, write, false)
.await?
.payload)
}
const BREAK_DRAIN_RECOVERY_TIMEOUT: Duration = Duration::from_secs(5);
async fn break_and_drain_wire(
read: &mut OracleReadHalf,
cx: &Cx,
write: &SharedWriteHalf,
recovery_timeout: Duration,
) -> Result<()> {
send_marker_shared(cx, write, TNS_MARKER_TYPE_BREAK)
.await
.map_err(|err| {
Error::ConnectionClosed(format!(
"failed to send break marker on call timeout: {err}"
))
})?;
match time::timeout(
time::wall_now(),
recovery_timeout,
drain_break_response(read, cx, write),
)
.await
{
Ok(Ok(())) => Ok(()),
Ok(Err(err)) => Err(Error::ConnectionClosed(format!(
"wire error while recovering from call timeout: {err}"
))),
Err(_) => Err(Error::ConnectionClosed(
"socket timed out while recovering from previous call timeout".to_string(),
)),
}
}
async fn cancel_and_drain_wire(
read: &mut OracleReadHalf,
cx: &Cx,
write: &SharedWriteHalf,
recovery_timeout: Duration,
) -> Result<()> {
break_and_drain_wire(read, cx, write, recovery_timeout).await
}
async fn drain_cancel_wire(
read: &mut OracleReadHalf,
cx: &Cx,
write: &SharedWriteHalf,
recovery_timeout: Duration,
) -> Result<()> {
match time::timeout(
time::wall_now(),
recovery_timeout,
drain_break_response(read, cx, write),
)
.await
{
Ok(Ok(())) => Ok(()),
Ok(Err(err)) => Err(Error::ConnectionClosed(format!(
"wire error while draining cancel response: {err}"
))),
Err(_) => Err(Error::ConnectionClosed(
"socket timed out while draining cancel response".to_string(),
)),
}
}
struct CancelDrainGuard<'a> {
pending: &'a Arc<AtomicBool>,
armed: bool,
}
impl<'a> CancelDrainGuard<'a> {
fn arm(pending: &'a Arc<AtomicBool>) -> Self {
Self {
pending,
armed: true,
}
}
fn disarm(&mut self) {
self.armed = false;
}
}
impl Drop for CancelDrainGuard<'_> {
fn drop(&mut self) {
if self.armed {
self.pending.store(true, Ordering::SeqCst);
}
}
}
async fn drain_break_response(
read: &mut OracleReadHalf,
cx: &Cx,
write: &SharedWriteHalf,
) -> Result<()> {
let initial_marker = loop {
let packet = read_packet(read, PacketLengthWidth::Large32).await?;
match packet.packet_type {
TNS_PACKET_TYPE_MARKER => break packet,
TNS_PACKET_TYPE_DATA => {
trace_connect_bytes("BREAK drain: discarded in-flight packet", &packet.payload);
continue;
}
other => {
return Err(oracledb_protocol::ProtocolError::UnknownMessageType {
message_type: other,
position: 4,
}
.into())
}
}
};
let pending = reset_after_marker(read, cx, write, &initial_marker).await?;
let trailing = read_data_response_boundary_from(read, cx, write, pending).await?;
trace_connect_bytes("BREAK drain: trailing error response", &trailing.payload);
Ok(())
}
async fn read_data_response_flushing_out_binds(
read: &mut OracleReadHalf,
cx: &Cx,
write: &SharedWriteHalf,
sdu: usize,
) -> Result<Vec<u8>> {
let mut response = read_data_response_boundary(read, cx, write, false).await?;
let mut payload = response.payload;
while response.flush_out_binds {
if matches!(payload.last(), Some(&TNS_MSG_TYPE_FLUSH_OUT_BINDS)) {
payload.pop();
}
send_data_packet_shared(cx, write, &[TNS_MSG_TYPE_FLUSH_OUT_BINDS], sdu).await?;
response = read_data_response_boundary(read, cx, write, false).await?;
payload.extend_from_slice(&response.payload);
}
Ok(payload)
}
fn data_packet_ends_response(flags: u16, payload: &[u8]) -> bool {
if flags
& (oracledb_protocol::thin::TNS_DATA_FLAGS_END_OF_RESPONSE
| oracledb_protocol::thin::TNS_DATA_FLAGS_EOF)
!= 0
{
return true;
}
payload == [TNS_MSG_TYPE_END_OF_RESPONSE] || payload == [TNS_MSG_TYPE_FLUSH_OUT_BINDS]
}
fn post_reset_packet_ends_response(payload: &[u8]) -> bool {
matches!(
payload.last(),
Some(&TNS_MSG_TYPE_FLUSH_OUT_BINDS) | Some(&TNS_MSG_TYPE_END_OF_RESPONSE)
)
}
async fn read_data_response_boundary(
read: &mut OracleReadHalf,
cx: &Cx,
write: &SharedWriteHalf,
in_pipeline: bool,
) -> Result<DataResponse> {
read_data_response_boundary_seeded(read, cx, write, in_pipeline, None).await
}
async fn read_data_response_boundary_from(
read: &mut OracleReadHalf,
cx: &Cx,
write: &SharedWriteHalf,
seed: Option<IncomingPacket>,
) -> Result<DataResponse> {
read_data_response_boundary_seeded(read, cx, write, false, seed).await
}
async fn read_data_response_boundary_seeded(
read: &mut OracleReadHalf,
cx: &Cx,
write: &SharedWriteHalf,
in_pipeline: bool,
seed: Option<IncomingPacket>,
) -> Result<DataResponse> {
let mut response = Vec::new();
let mut pending_packet = seed;
let mut after_reset = false;
loop {
let packet = match pending_packet.take() {
Some(packet) => packet,
None => read_packet(read, PacketLengthWidth::Large32).await?,
};
if packet.packet_type == TNS_PACKET_TYPE_MARKER {
if in_pipeline {
trace_connect_bytes("MARKER packet skipped in pipeline", &packet.payload);
continue;
}
pending_packet = reset_after_marker(read, cx, write, &packet).await?;
after_reset = true;
continue;
}
if packet.packet_type != TNS_PACKET_TYPE_DATA {
return Err(oracledb_protocol::ProtocolError::UnknownMessageType {
message_type: packet.packet_type,
position: 4,
}
.into());
}
let (data_flags, payload) = packet.payload.split_at_checked(2).ok_or(
oracledb_protocol::ProtocolError::TtcDecode("missing data packet flags"),
)?;
let flags = u16::from_be_bytes(
data_flags
.try_into()
.map_err(|_| oracledb_protocol::ProtocolError::TtcDecode("invalid flags"))?,
);
let ends = data_packet_ends_response(flags, payload)
|| (after_reset && post_reset_packet_ends_response(payload));
if ends && response.is_empty() {
response = packet.payload;
response.drain(..2);
break;
}
response.extend_from_slice(payload);
if ends {
break;
}
}
let flush_out_binds = matches!(response.last(), Some(&TNS_MSG_TYPE_FLUSH_OUT_BINDS));
Ok(DataResponse {
payload: response,
flush_out_binds,
})
}
const TNS_PACKET_TYPE_MARKER: u8 = 12;
const TNS_MARKER_TYPE_BREAK: u8 = 1;
const TNS_MARKER_TYPE_RESET: u8 = 2;
async fn reset_after_marker(
read: &mut OracleReadHalf,
cx: &Cx,
write: &SharedWriteHalf,
initial_marker: &IncomingPacket,
) -> Result<Option<IncomingPacket>> {
trace_connect_bytes("MARKER packet", &initial_marker.payload);
send_marker_shared(cx, write, TNS_MARKER_TYPE_RESET).await?;
loop {
let packet = read_packet(read, PacketLengthWidth::Large32).await?;
if packet.packet_type != TNS_PACKET_TYPE_MARKER {
return Ok(Some(packet));
}
trace_connect_bytes("MARKER reset response", &packet.payload);
}
}
async fn send_marker<W>(stream: &mut W, marker_type: u8) -> Result<()>
where
W: AsyncWrite + Unpin,
{
let packet = encode_packet(
TNS_PACKET_TYPE_MARKER,
0,
None,
&[1, 0, marker_type],
PacketLengthWidth::Large32,
)?;
trace_connect_bytes("send MARKER", &packet);
stream.write_all(&packet).await?;
stream.flush().await?;
Ok(())
}
async fn read_packet<R>(stream: &mut R, width: PacketLengthWidth) -> Result<IncomingPacket>
where
R: AsyncRead + Unpin,
{
let mut header = [0u8; 8];
stream.read_exact(&mut header).await?;
let [len0, len1, len2, len3, packet_type, _, _, _] = header;
let declared = match width {
PacketLengthWidth::Legacy16 => usize::from(u16::from_be_bytes([len0, len1])),
PacketLengthWidth::Large32 => {
usize::try_from(u32::from_be_bytes([len0, len1, len2, len3])).unwrap_or(usize::MAX)
}
};
if declared < header.len() {
return Err(oracledb_protocol::ProtocolError::InvalidPacketLength {
length: declared,
minimum: header.len(),
}
.into());
}
let mut payload = vec![0u8; declared - header.len()];
stream.read_exact(&mut payload).await?;
Ok(IncomingPacket {
packet_type,
payload,
})
}
fn listener_connect_descriptor_with_server(
descriptor: &EasyConnect,
identity: &ClientIdentity,
server_type_emon: bool,
) -> String {
let server = if server_type_emon {
"(SERVER=emon)"
} else {
""
};
format!(
"(DESCRIPTION=(ADDRESS=(PROTOCOL=tcp)(HOST={})(PORT={}))(CONNECT_DATA=(SERVICE_NAME={}){}(CID=(PROGRAM={})(HOST={})(USER={}))))",
descriptor.host,
descriptor.port,
descriptor.service_name,
server,
identity.program,
identity.machine,
identity.osuser,
)
}
fn auth_connect_descriptor(descriptor: &EasyConnect) -> String {
format!(
"(DESCRIPTION=(ADDRESS=(PROTOCOL=tcp)(HOST={})(PORT={}))(CONNECT_DATA=(SERVICE_NAME={})))",
descriptor.host, descriptor.port, descriptor.service_name
)
}
fn parse_session_u32(
data: &std::collections::BTreeMap<String, String>,
key: &'static str,
) -> Result<u32> {
data.get(key)
.ok_or(Error::MissingSessionField(key))?
.parse::<u32>()
.map_err(|_| Error::MissingSessionField(key))
}
fn parse_session_u16(
data: &std::collections::BTreeMap<String, String>,
key: &'static str,
) -> Result<u16> {
data.get(key)
.ok_or(Error::MissingSessionField(key))?
.parse::<u16>()
.map_err(|_| Error::MissingSessionField(key))
}
fn next_ttc_sequence(seq_num: &mut u8) -> u8 {
*seq_num = seq_num.wrapping_add(1);
if *seq_num == 0 {
*seq_num = 1;
}
*seq_num
}
fn statement_is_query(sql: &str) -> bool {
sql.trim_start()
.split(|ch: char| !ch.is_ascii_alphabetic())
.next()
.is_some_and(|keyword| keyword.eq_ignore_ascii_case("select"))
}
fn columns_require_define(columns: &[ColumnMetadata]) -> bool {
use oracledb_protocol::thin::{
ORA_TYPE_NUM_BLOB, ORA_TYPE_NUM_CLOB, ORA_TYPE_NUM_JSON, ORA_TYPE_NUM_VECTOR,
};
columns.iter().any(|column| {
matches!(
column.ora_type_num,
ORA_TYPE_NUM_CLOB | ORA_TYPE_NUM_BLOB | ORA_TYPE_NUM_VECTOR | ORA_TYPE_NUM_JSON
)
})
}
fn json_lob_probe_candidates(columns: &[ColumnMetadata]) -> Vec<(usize, String)> {
use oracledb_protocol::thin::{ORA_TYPE_NUM_BLOB, ORA_TYPE_NUM_CLOB};
columns
.iter()
.enumerate()
.filter(|(_, metadata)| {
!metadata.is_json
&& matches!(metadata.ora_type_num, ORA_TYPE_NUM_CLOB | ORA_TYPE_NUM_BLOB)
&& !metadata.name.is_empty()
})
.map(|(index, metadata)| (index, metadata.name.to_ascii_uppercase()))
.collect()
}
fn trace_connect_step(step: &'static str) {
if std::env::var_os("ORACLEDB_TRACE_CONNECT").is_some() {
eprintln!("oracledb::connect: {step}");
}
}
fn trace_connect_value(label: &'static str, value: &str) {
if std::env::var_os("ORACLEDB_TRACE_CONNECT").is_some() {
eprintln!("oracledb::connect: {label}: {value}");
}
}
fn trace_connect_bytes(label: &'static str, bytes: &[u8]) {
if std::env::var_os("ORACLEDB_TRACE_CONNECT").is_some() {
let mut hex = String::with_capacity(bytes.len() * 2);
for byte in bytes {
use std::fmt::Write as _;
let _ = write!(&mut hex, "{byte:02x}");
}
eprintln!("oracledb::connect: {label} len={} hex={hex}", bytes.len());
}
}
fn trace_query_bytes(label: &'static str, bytes: &[u8]) {
if std::env::var_os("ORACLEDB_TRACE_QUERY").is_some() {
let mut hex = String::with_capacity(bytes.len() * 2);
for byte in bytes {
use std::fmt::Write as _;
let _ = write!(&mut hex, "{byte:02x}");
}
eprintln!("oracledb::query: {label} len={} hex={hex}", bytes.len());
}
}
#[cfg(test)]
mod tests {
use super::*;
use std::io::Read;
use std::net::TcpListener;
use std::thread;
use std::time::Duration;
fn identity() -> ClientIdentity {
ClientIdentity::new("program", "machine", "osuser", "terminal", "driver")
.expect("test identity should be valid")
}
fn server_error(message: &str) -> Error {
Error::Protocol(oracledb_protocol::ProtocolError::ServerError(
message.to_string(),
))
}
fn structured_error(code: u32, pos: i32) -> Error {
Error::Protocol(oracledb_protocol::ProtocolError::ServerErrorInfo(Box::new(
oracledb_protocol::ServerErrorDetails {
message: format!("ORA-{code:05}: synthetic"),
code,
pos,
..Default::default()
},
)))
}
fn column(name: &str, ora_type_num: u8, is_json: bool) -> ColumnMetadata {
ColumnMetadata {
name: name.to_string(),
ora_type_num,
is_json,
..ColumnMetadata::default()
}
}
#[test]
fn json_lob_probe_candidates_selects_named_non_json_lobs() {
use oracledb_protocol::thin::{ORA_TYPE_NUM_BLOB, ORA_TYPE_NUM_CLOB, ORA_TYPE_NUM_VARCHAR};
let columns = vec![
column("doc", ORA_TYPE_NUM_CLOB, false), column("blob_doc", ORA_TYPE_NUM_BLOB, false), column("already_json", ORA_TYPE_NUM_CLOB, true), column("name", ORA_TYPE_NUM_VARCHAR, false), column("", ORA_TYPE_NUM_CLOB, false), ];
assert_eq!(
json_lob_probe_candidates(&columns),
vec![(0, "DOC".to_string()), (1, "BLOB_DOC".to_string())]
);
}
#[test]
fn json_lob_probe_candidates_empty_when_no_lobs() {
use oracledb_protocol::thin::ORA_TYPE_NUM_VARCHAR;
let columns = vec![column("name", ORA_TYPE_NUM_VARCHAR, false)];
assert!(json_lob_probe_candidates(&columns).is_empty());
}
#[test]
fn ora_code_parses_from_message_and_struct() {
assert_eq!(
server_error("ORA-00060: deadlock detected").ora_code(),
Some(60)
);
assert_eq!(structured_error(942, 0).ora_code(), Some(942));
assert_eq!(server_error("listener problem").ora_code(), None);
assert_eq!(Error::CallTimeout(500).ora_code(), None);
}
#[test]
fn offset_only_from_structured_nonzero() {
assert_eq!(structured_error(942, 14).offset(), Some(14));
assert_eq!(structured_error(942, 0).offset(), None);
assert_eq!(
server_error("ORA-00942: table or view does not exist").offset(),
None
);
}
#[test]
fn transient_classification() {
for &code in TRANSIENT_ORA_CODES {
let err = server_error(&format!("ORA-{code:05}: transient"));
assert!(err.is_transient(), "ORA-{code:05} should be transient");
assert!(err.is_retryable(), "transient implies retryable");
assert!(
!err.is_connection_lost(),
"ORA-{code:05} is not connection-lost"
);
}
let perm = server_error("ORA-00942: table or view does not exist");
assert!(!perm.is_transient());
assert!(!perm.is_connection_lost());
assert!(!perm.is_retryable());
}
#[test]
fn connection_lost_classification() {
for &code in CONNECTION_LOST_ORA_CODES {
let err = server_error(&format!("ORA-{code:05}: lost"));
assert!(
err.is_connection_lost(),
"ORA-{code:05} should be connection-lost"
);
assert!(err.is_retryable(), "connection-lost implies retryable");
assert!(
!err.is_transient(),
"ORA-{code:05} is not a transient (retry-in-place) code"
);
}
let io = Error::Io(std::io::Error::new(
std::io::ErrorKind::ConnectionReset,
"reset",
));
assert!(io.is_connection_lost());
assert!(io.is_retryable());
let timeout = Error::CallTimeout(1000);
assert!(
!timeout.is_connection_lost(),
"a call timeout leaves the connection usable after the drain"
);
assert!(
timeout.is_transient(),
"a call timeout is a retry-in-place (transient) condition"
);
assert!(
timeout.is_retryable(),
"transient implies retryable on the same connection"
);
let recovery_failed =
Error::ConnectionClosed("socket timed out while recovering".to_string());
assert!(
recovery_failed.is_connection_lost(),
"a failed timeout-recovery drain marks the connection lost"
);
assert!(recovery_failed.is_retryable(), "reconnect, then retry");
assert!(
!recovery_failed.is_transient(),
"ConnectionClosed needs a reconnect first, so it is not retry-in-place"
);
}
#[test]
fn data_packet_ends_response_requires_flag_or_lone_marker_byte() {
const EOR: u8 = TNS_MSG_TYPE_END_OF_RESPONSE; const FOB: u8 = TNS_MSG_TYPE_FLUSH_OUT_BINDS; let eor_flag = oracledb_protocol::thin::TNS_DATA_FLAGS_END_OF_RESPONSE;
let eof_flag = oracledb_protocol::thin::TNS_DATA_FLAGS_EOF;
assert!(data_packet_ends_response(eor_flag, &[0x01, 0x02, EOR]));
assert!(data_packet_ends_response(eor_flag, &[]));
assert!(data_packet_ends_response(eof_flag, &[0x01, 0x02, 0x03]));
assert!(data_packet_ends_response(0x0000, &[EOR]));
assert!(data_packet_ends_response(0x0000, &[FOB]));
assert!(!data_packet_ends_response(0x0000, &[0xc1, 0x02, EOR]));
assert!(!data_packet_ends_response(0x0000, &[0x00, EOR]));
assert!(!data_packet_ends_response(0x0000, &[EOR, 0x05, 0x06, EOR]));
assert!(!data_packet_ends_response(0x0000, &[0xc1, 0x02, FOB]));
assert!(!data_packet_ends_response(0x0000, &[0x00, FOB]));
assert!(!data_packet_ends_response(0x0000, &[0x01, 0x02, 0x03]));
assert!(!data_packet_ends_response(0x0000, &[]));
}
fn replay_boundary(packets: &[(u16, Vec<u8>)]) -> (Vec<u8>, Option<usize>, bool) {
let mut reassembled = Vec::new();
let mut stopped_at = None;
for (index, (flags, payload)) in packets.iter().enumerate() {
let ends = data_packet_ends_response(*flags, payload);
if ends && reassembled.is_empty() {
reassembled = payload.clone();
stopped_at = Some(index);
break;
}
reassembled.extend_from_slice(payload);
if ends {
stopped_at = Some(index);
break;
}
}
let flush_out_binds = matches!(reassembled.last(), Some(&TNS_MSG_TYPE_FLUSH_OUT_BINDS));
(reassembled, stopped_at, flush_out_binds)
}
#[test]
fn boundary_loop_reassembles_packets_ending_in_marker_byte() {
const EOR: u8 = TNS_MSG_TYPE_END_OF_RESPONSE;
const FOB: u8 = TNS_MSG_TYPE_FLUSH_OUT_BINDS;
let packets: [(u16, Vec<u8>); 5] = [
(0x0000, vec![0x10, 0x11, EOR]), (0x0000, vec![0x20, 0x21, 0x22, FOB]), (0x0000, vec![0x30, 0x31, 0x32, EOR]), (0x0000, vec![0x33, 0x34, 0x35]), (
oracledb_protocol::thin::TNS_DATA_FLAGS_END_OF_RESPONSE,
vec![0x40, 0x41, EOR], ),
];
let (reassembled, stopped_at, flush_out_binds) = replay_boundary(&packets);
assert_eq!(
stopped_at,
Some(4),
"reassembly must stop only on the flagged final packet, not on a body packet ending in a marker byte"
);
assert!(
!flush_out_binds,
"the response ended in END_OF_RESPONSE, not FLUSH_OUT_BINDS"
);
assert_eq!(
reassembled,
vec![
0x10, 0x11, EOR, 0x20, 0x21, 0x22, FOB, 0x30, 0x31, 0x32, EOR, 0x33, 0x34, 0x35, 0x40, 0x41, EOR, ],
"every body packet's bytes must be concatenated in order with none dropped"
);
}
#[test]
fn boundary_loop_detects_flush_out_binds_only_at_true_boundary() {
const FOB: u8 = TNS_MSG_TYPE_FLUSH_OUT_BINDS;
let packets: [(u16, Vec<u8>); 2] = [
(0x0000, vec![0x01, 0x02, FOB]), (
oracledb_protocol::thin::TNS_DATA_FLAGS_END_OF_RESPONSE,
vec![0x03, FOB], ),
];
let (reassembled, stopped_at, flush_out_binds) = replay_boundary(&packets);
assert_eq!(stopped_at, Some(1), "stop on the EOR-flagged tail");
assert!(
flush_out_binds,
"flush-out-binds must be detected from the terminal FLUSH_OUT_BINDS message byte"
);
assert_eq!(reassembled, vec![0x01, 0x02, FOB, 0x03, FOB]);
}
#[test]
fn single_packet_passthrough_is_byte_identical_to_extend() {
const EOR: u8 = TNS_MSG_TYPE_END_OF_RESPONSE;
const FOB: u8 = TNS_MSG_TYPE_FLUSH_OUT_BINDS;
let eor_flag = oracledb_protocol::thin::TNS_DATA_FLAGS_END_OF_RESPONSE;
fn replay_extend_only(packets: &[(u16, Vec<u8>)]) -> (Vec<u8>, Option<usize>, bool) {
let mut reassembled = Vec::new();
let mut stopped_at = None;
for (index, (flags, payload)) in packets.iter().enumerate() {
reassembled.extend_from_slice(payload);
if data_packet_ends_response(*flags, payload) {
stopped_at = Some(index);
break;
}
}
let flush = matches!(reassembled.last(), Some(&TNS_MSG_TYPE_FLUSH_OUT_BINDS));
(reassembled, stopped_at, flush)
}
let cases: &[(u16, Vec<u8>)] = &[
(eor_flag, vec![0x40, 0x41, EOR]),
(eor_flag, vec![0x03, FOB]), (eor_flag, vec![0xde, 0xad, 0xbe, 0xef]),
(eor_flag, vec![0x00]),
];
for (flags, payload) in cases {
let one = [(*flags, payload.clone())];
let passthrough = replay_boundary(&one);
let extend = replay_extend_only(&one);
assert_eq!(
passthrough, extend,
"passthrough must equal extend for single packet {payload:02x?}"
);
assert_eq!(&passthrough.0, payload);
}
}
#[test]
fn descriptor_builder_uses_identity_in_listener_cid() {
let options = ConnectOptions::new("localhost/FREEPDB1", "user", "password", identity());
let descriptor =
EasyConnect::parse(&options.connect_string).expect("test connect string should parse");
let built = listener_connect_descriptor_with_server(&descriptor, &options.identity, false);
assert!(built.contains("(PROGRAM=program)"));
assert!(built.contains("(HOST=machine)"));
assert!(built.contains("(USER=osuser)"));
assert!(!built.contains("(SERVER=emon)"));
let emon = listener_connect_descriptor_with_server(&descriptor, &options.identity, true);
assert!(emon.contains("(SERVICE_NAME=FREEPDB1)(SERVER=emon)(CID="));
}
#[test]
fn cancel_handle_sends_tns_break_marker() {
let listener = TcpListener::bind("127.0.0.1:0").expect("bind local listener");
let addr = listener.local_addr().expect("listener address");
let server = thread::spawn(move || {
let (mut socket, _) = listener.accept().expect("accept test client");
socket
.set_read_timeout(Some(Duration::from_secs(5)))
.expect("set read timeout");
let mut packet = [0u8; 11];
socket.read_exact(&mut packet).expect("read marker packet");
packet
});
let runtime = build_io_runtime().expect("asupersync runtime");
let mut handle = runtime.block_on(async {
let stream = TcpStream::connect(addr).await.expect("connect to listener");
let (_read, write) = transport::plain_split(stream);
CancelHandle {
write: Arc::new(AsyncMutex::with_name("oracle_tcp_write_test", write)),
}
});
handle.cancel().expect("cancel marker write");
let packet = server.join().expect("server thread joins");
assert_eq!(
packet,
[
0,
0,
0,
11,
TNS_PACKET_TYPE_MARKER,
0,
0,
0,
1,
0,
TNS_MARKER_TYPE_BREAK
]
);
}
const EOR_FLAG: u16 = oracledb_protocol::thin::TNS_DATA_FLAGS_END_OF_RESPONSE;
fn data_packet(message: &[u8], end_of_response: bool) -> Vec<u8> {
let flags = if end_of_response { EOR_FLAG } else { 0 };
encode_packet(
TNS_PACKET_TYPE_DATA,
0,
Some(flags),
message,
PacketLengthWidth::Large32,
)
.expect("encode data packet")
}
fn marker_packet(marker_type: u8) -> Vec<u8> {
encode_packet(
TNS_PACKET_TYPE_MARKER,
0,
None,
&[1, 0, marker_type],
PacketLengthWidth::Large32,
)
.expect("encode marker packet")
}
fn read_marker_type(socket: &mut std::net::TcpStream) -> u8 {
let mut packet = [0u8; 11];
socket.read_exact(&mut packet).expect("read marker packet");
assert_eq!(
packet[4], TNS_PACKET_TYPE_MARKER,
"expected a MARKER packet"
);
packet[10]
}
#[test]
fn break_and_drain_consumes_inflight_response_and_reset_then_next_read_is_fresh() {
const INFLIGHT_BODY: &[u8] = &[0xDE, 0xAD, 0xBE, 0xEF];
const ERROR_BODY: &[u8] = &[0x04, 0x01, 0x02];
const FRESH_BODY: &[u8] = &[0x11, 0x22, 0x33, 0x44, 0x55];
let listener = TcpListener::bind("127.0.0.1:0").expect("bind local listener");
let addr = listener.local_addr().expect("listener address");
let server = thread::spawn(move || {
let (mut socket, _) = listener.accept().expect("accept test client");
socket
.set_read_timeout(Some(Duration::from_secs(5)))
.expect("set read timeout");
use std::io::Write as _;
assert_eq!(
read_marker_type(&mut socket),
TNS_MARKER_TYPE_BREAK,
"client must send a BREAK marker first"
);
socket
.write_all(&data_packet(INFLIGHT_BODY, true))
.expect("write in-flight response");
socket
.write_all(&marker_packet(TNS_MARKER_TYPE_BREAK))
.expect("write break-ack marker");
assert_eq!(
read_marker_type(&mut socket),
TNS_MARKER_TYPE_RESET,
"client must answer the marker with a RESET"
);
socket
.write_all(&marker_packet(TNS_MARKER_TYPE_RESET))
.expect("write reset-confirm marker");
socket
.write_all(&data_packet(ERROR_BODY, true))
.expect("write trailing error packet");
socket
.write_all(&data_packet(FRESH_BODY, true))
.expect("write fresh response");
});
let runtime = build_io_runtime().expect("asupersync runtime");
let next = runtime.block_on(async {
let cx = Cx::current().expect("ambient Cx");
let stream = TcpStream::connect(addr).await.expect("connect to listener");
let (mut read, write) = transport::plain_split(stream);
let write: SharedWriteHalf = Arc::new(AsyncMutex::with_name("drain_test_write", write));
break_and_drain_wire(&mut read, &cx, &write, Duration::from_secs(5))
.await
.expect("drain must succeed and leave the stream clean");
read_data_response(&mut read, &cx, &write)
.await
.expect("next read after drain must decode cleanly")
});
assert_eq!(
next, FRESH_BODY,
"after break_and_drain the reused connection must read the FRESH response, \
not the stale in-flight response ({INFLIGHT_BODY:?}) or error body ({ERROR_BODY:?})"
);
server.join().expect("server thread joins");
}
#[test]
fn reset_after_marker_drains_multiple_trailing_markers_no_duplicate_reset() {
const INFLIGHT_BODY: &[u8] = &[0xDE, 0xAD];
const ERROR_BODY: &[u8] = &[0x04, 0x01, 0x02];
const FRESH_BODY: &[u8] = &[0x11, 0x22, 0x33];
let listener = TcpListener::bind("127.0.0.1:0").expect("bind local listener");
let addr = listener.local_addr().expect("listener address");
let server = thread::spawn(move || {
let (mut socket, _) = listener.accept().expect("accept test client");
socket
.set_read_timeout(Some(Duration::from_secs(5)))
.expect("set read timeout");
use std::io::Write as _;
assert_eq!(
read_marker_type(&mut socket),
TNS_MARKER_TYPE_BREAK,
"client must send a BREAK marker first"
);
socket
.write_all(&data_packet(INFLIGHT_BODY, true))
.expect("write in-flight response");
socket
.write_all(&marker_packet(TNS_MARKER_TYPE_BREAK))
.expect("write break-ack marker");
assert_eq!(
read_marker_type(&mut socket),
TNS_MARKER_TYPE_RESET,
"client must answer with a RESET"
);
socket
.write_all(&marker_packet(TNS_MARKER_TYPE_RESET))
.expect("write reset marker #1");
socket
.write_all(&marker_packet(TNS_MARKER_TYPE_RESET))
.expect("write reset marker #2");
socket
.write_all(&data_packet(ERROR_BODY, true))
.expect("write trailing error packet");
socket
.write_all(&data_packet(FRESH_BODY, true))
.expect("write fresh response");
socket
.set_read_timeout(Some(Duration::from_millis(750)))
.expect("set short read timeout");
let mut extra = [0u8; 11];
if socket.read_exact(&mut extra).is_ok() {
panic!(
"client sent a DUPLICATE marker (type {}): the drain did not \
consume all trailing RESET markers (bead rust-oracledb-yhz)",
extra[10]
);
}
});
let runtime = build_io_runtime().expect("asupersync runtime");
let next = runtime.block_on(async {
let cx = Cx::current().expect("ambient Cx");
let stream = TcpStream::connect(addr).await.expect("connect to listener");
let (mut read, write) = transport::plain_split(stream);
let write: SharedWriteHalf = Arc::new(AsyncMutex::with_name("yhz_test_write", write));
break_and_drain_wire(&mut read, &cx, &write, Duration::from_secs(5))
.await
.expect("drain must succeed even with multiple RESET markers");
read_data_response(&mut read, &cx, &write)
.await
.expect("next read after drain must decode cleanly")
});
assert_eq!(
next, FRESH_BODY,
"after draining multiple RESET markers the reused connection must read \
the FRESH response"
);
server.join().expect("server thread joins");
}
#[test]
fn break_without_drain_leaves_stale_bytes_for_next_read() {
const STALE_BODY: &[u8] = &[0x53, 0x54, 0x41, 0x4c, 0x45]; const FRESH_BODY: &[u8] = &[0x11, 0x22, 0x33];
let listener = TcpListener::bind("127.0.0.1:0").expect("bind local listener");
let addr = listener.local_addr().expect("listener address");
let server = thread::spawn(move || {
let (mut socket, _) = listener.accept().expect("accept test client");
socket
.set_read_timeout(Some(Duration::from_secs(5)))
.expect("set read timeout");
use std::io::Write as _;
assert_eq!(read_marker_type(&mut socket), TNS_MARKER_TYPE_BREAK);
socket
.write_all(&data_packet(STALE_BODY, true))
.expect("write stale in-flight response");
socket
.write_all(&data_packet(FRESH_BODY, true))
.expect("write fresh response");
});
let runtime = build_io_runtime().expect("asupersync runtime");
let first_read = runtime.block_on(async {
let cx = Cx::current().expect("ambient Cx");
let stream = TcpStream::connect(addr).await.expect("connect to listener");
let (mut read, write) = transport::plain_split(stream);
let write: SharedWriteHalf =
Arc::new(AsyncMutex::with_name("nodrain_test_write", write));
send_marker_shared(&cx, &write, TNS_MARKER_TYPE_BREAK)
.await
.expect("send break");
read_data_response(&mut read, &cx, &write)
.await
.expect("read after bare break")
});
assert_eq!(
first_read, STALE_BODY,
"without the drain, the next read misframes onto the stale in-flight bytes — \
this is the bug break_and_drain fixes"
);
server.join().expect("server thread joins");
}
#[test]
fn dml_returning_error_flush_out_binds_after_reset_completes_without_hang() {
const ERROR_BODY: &[u8] = &[0x04, 0x01, 0x02, 0x37];
const FLUSH_REQUEST_BODY: &[u8] = &[0x07, 0x00, 0x00, TNS_MSG_TYPE_FLUSH_OUT_BINDS];
let listener = TcpListener::bind("127.0.0.1:0").expect("bind local listener");
let addr = listener.local_addr().expect("listener address");
let server = thread::spawn(move || {
let (mut socket, _) = listener.accept().expect("accept test client");
socket
.set_read_timeout(Some(Duration::from_secs(5)))
.expect("set read timeout");
use std::io::Write as _;
socket
.write_all(&marker_packet(TNS_MARKER_TYPE_BREAK))
.expect("write break marker");
assert_eq!(
read_marker_type(&mut socket),
TNS_MARKER_TYPE_RESET,
"client must answer the BREAK with a RESET"
);
socket
.write_all(&marker_packet(TNS_MARKER_TYPE_RESET))
.expect("write reset-confirm marker");
socket
.write_all(&data_packet(FLUSH_REQUEST_BODY, false))
.expect("write flush-out-binds request");
let mut header = [0u8; 8];
socket
.read_exact(&mut header)
.expect("read flush-out-binds reply header");
assert_eq!(
header[4], TNS_PACKET_TYPE_DATA,
"client's flush-out-binds reply must be a DATA packet"
);
let len = u32::from_be_bytes([header[0], header[1], header[2], header[3]]) as usize;
let mut body = vec![0u8; len - 8];
socket
.read_exact(&mut body)
.expect("read flush-out-binds reply body");
assert_eq!(
body.last().copied(),
Some(TNS_MSG_TYPE_FLUSH_OUT_BINDS),
"client must reply with a FLUSH_OUT_BINDS message"
);
socket
.write_all(&marker_packet(TNS_MARKER_TYPE_BREAK))
.expect("write second break marker");
assert_eq!(
read_marker_type(&mut socket),
TNS_MARKER_TYPE_RESET,
"client must answer the second BREAK with a RESET"
);
socket
.write_all(&marker_packet(TNS_MARKER_TYPE_RESET))
.expect("write second reset-confirm marker");
socket
.write_all(&data_packet(ERROR_BODY, true))
.expect("write trailing ORA-12899 error packet");
});
let runtime = build_io_runtime().expect("asupersync runtime");
let payload = runtime.block_on(async {
let cx = Cx::current().expect("ambient Cx");
let stream = TcpStream::connect(addr).await.expect("connect to listener");
let (mut read, write) = transport::plain_split(stream);
let write: SharedWriteHalf =
Arc::new(AsyncMutex::with_name("returning_err_test_write", write));
time::timeout(
time::wall_now(),
Duration::from_secs(10),
read_data_response_flushing_out_binds(&mut read, &cx, &write, 8192),
)
.await
.expect("must NOT hang on the DML-RETURNING error path (flush-out-binds after reset)")
.expect("read must complete and yield the trailing error payload")
});
assert!(
payload.ends_with(ERROR_BODY),
"the reassembled response must end with the ORA-12899 error payload, got {payload:?}"
);
server.join().expect("server thread joins");
}
#[test]
fn cancel_and_drain_wire_leaves_connection_reusable() {
const INFLIGHT_BODY: &[u8] = &[0xCA, 0xFE];
const ERROR_BODY: &[u8] = &[0x04, 0x01, 0x0d];
const FRESH_BODY: &[u8] = &[0x07, 0x05, 0x0c];
let listener = TcpListener::bind("127.0.0.1:0").expect("bind local listener");
let addr = listener.local_addr().expect("listener address");
let server = thread::spawn(move || {
let (mut socket, _) = listener.accept().expect("accept test client");
socket
.set_read_timeout(Some(Duration::from_secs(5)))
.expect("set read timeout");
use std::io::Write as _;
assert_eq!(
read_marker_type(&mut socket),
TNS_MARKER_TYPE_BREAK,
"cancel must send a BREAK marker first"
);
socket
.write_all(&data_packet(INFLIGHT_BODY, true))
.expect("write in-flight response");
socket
.write_all(&marker_packet(TNS_MARKER_TYPE_BREAK))
.expect("write break-ack marker");
assert_eq!(
read_marker_type(&mut socket),
TNS_MARKER_TYPE_RESET,
"cancel must answer the marker with a RESET"
);
socket
.write_all(&marker_packet(TNS_MARKER_TYPE_RESET))
.expect("write reset-confirm marker");
socket
.write_all(&data_packet(ERROR_BODY, true))
.expect("write trailing error packet");
socket
.write_all(&data_packet(FRESH_BODY, true))
.expect("write fresh response");
});
let runtime = build_io_runtime().expect("asupersync runtime");
let next = runtime.block_on(async {
let cx = Cx::current().expect("ambient Cx");
let stream = TcpStream::connect(addr).await.expect("connect to listener");
let (mut read, write) = transport::plain_split(stream);
let write: SharedWriteHalf =
Arc::new(AsyncMutex::with_name("cancel_test_write", write));
cancel_and_drain_wire(&mut read, &cx, &write, Duration::from_secs(5))
.await
.expect("cancel drain must succeed and leave the stream clean");
read_data_response(&mut read, &cx, &write)
.await
.expect("next read after cancel must decode cleanly")
});
assert_eq!(
next, FRESH_BODY,
"after cancel the reused connection must read the FRESH response, not the \
stale in-flight response ({INFLIGHT_BODY:?}) or error body ({ERROR_BODY:?})"
);
server.join().expect("server thread joins");
}
#[test]
fn drain_cancel_wire_drains_without_sending_a_break() {
const INFLIGHT_BODY: &[u8] = &[0xCA, 0xFE];
const ERROR_BODY: &[u8] = &[0x04, 0x01, 0x0d];
const FRESH_BODY: &[u8] = &[0x07, 0x05, 0x0c];
let listener = TcpListener::bind("127.0.0.1:0").expect("bind local listener");
let addr = listener.local_addr().expect("listener address");
let server = thread::spawn(move || {
let (mut socket, _) = listener.accept().expect("accept test client");
socket
.set_read_timeout(Some(Duration::from_secs(5)))
.expect("set read timeout");
use std::io::Write as _;
socket
.write_all(&data_packet(INFLIGHT_BODY, true))
.expect("write in-flight response");
socket
.write_all(&marker_packet(TNS_MARKER_TYPE_BREAK))
.expect("write break-ack marker");
assert_eq!(
read_marker_type(&mut socket),
TNS_MARKER_TYPE_RESET,
"drain must answer the break-ack marker with a RESET"
);
socket
.write_all(&marker_packet(TNS_MARKER_TYPE_RESET))
.expect("write reset-confirm marker");
socket
.write_all(&data_packet(ERROR_BODY, true))
.expect("write trailing error packet");
socket
.write_all(&data_packet(FRESH_BODY, true))
.expect("write fresh response");
socket
.set_read_timeout(Some(Duration::from_millis(500)))
.expect("set short read timeout");
let mut extra = [0u8; 11];
if let Ok(()) = socket.read_exact(&mut extra) {
assert_ne!(
extra[10], TNS_MARKER_TYPE_BREAK,
"drain-only cancel must NOT send a BREAK marker"
);
}
});
let runtime = build_io_runtime().expect("asupersync runtime");
let next = runtime.block_on(async {
let cx = Cx::current().expect("ambient Cx");
let stream = TcpStream::connect(addr).await.expect("connect to listener");
let (mut read, write) = transport::plain_split(stream);
let write: SharedWriteHalf =
Arc::new(AsyncMutex::with_name("drain_cancel_test_write", write));
drain_cancel_wire(&mut read, &cx, &write, Duration::from_secs(5))
.await
.expect("drain-only cancel must succeed and leave the stream clean");
read_data_response(&mut read, &cx, &write)
.await
.expect("next read after drain must decode cleanly")
});
assert_eq!(
next, FRESH_BODY,
"after drain-only cancel the reused connection must read the FRESH response"
);
server.join().expect("server thread joins");
}
#[test]
fn cancel_drain_guard_arms_pending_flag_only_when_dropped_in_flight() {
let pending = Arc::new(AtomicBool::new(false));
{
let _guard = CancelDrainGuard::arm(&pending);
}
assert!(
pending.load(Ordering::SeqCst),
"dropping an armed guard (cancelled in flight) must arm the drain flag"
);
pending.store(false, Ordering::SeqCst);
{
let mut guard = CancelDrainGuard::arm(&pending);
guard.disarm();
}
assert!(
!pending.load(Ordering::SeqCst),
"a disarmed guard (clean completion) must NOT arm the drain flag"
);
}
#[test]
fn cancelled_error_is_not_connection_lost_but_is_transient() {
let cancelled = Error::Cancelled;
assert!(
!cancelled.is_connection_lost(),
"a user cancel leaves the session alive (ORA-01013 / DPY-4024 semantics)"
);
assert!(
cancelled.is_transient(),
"a cancelled operation may be retried on the same clean connection"
);
assert_eq!(cancelled.ora_code(), Some(1013));
}
}