use super::character_set::CharacterSet;
use super::encode_int::encode_length_int;
use super::new_writer::NewWriter;
use crate::mysql::capabilities::Capabilities;
use crate::shared::data::Data;
use bun_collections::StringHashMap;
bun_core::declare_scope!(MySQLConnection, hidden);
pub struct HandshakeResponse41 {
pub capability_flags: Capabilities,
pub max_packet_size: u32, pub character_set: CharacterSet, pub username: Data,
pub auth_response: Data,
pub database: Data,
pub auth_plugin_name: Data,
pub connect_attrs: StringHashMap<Box<[u8]>>, pub sequence_id: u8,
}
impl HandshakeResponse41 {
pub fn write_internal<Context: super::new_writer::WriterContext>(
&mut self,
writer: NewWriter<Context>,
) -> Result<(), bun_core::Error> {
let mut packet = writer.start(self.sequence_id)?;
self.capability_flags.CLIENT_CONNECT_ATTRS = self.connect_attrs.len() > 0;
let caps = self.capability_flags.to_int();
writer.int4(caps)?;
bun_core::scoped_log!(
MySQLConnection,
"Client capabilities: [{}] 0x{:08x} sequence_id: {}",
self.capability_flags,
caps,
self.sequence_id
);
writer.int4(self.max_packet_size)?;
writer.int1(self.character_set as u8)?;
writer.write(&[0u8; 23])?;
writer.write_z(self.username.slice())?;
let auth_data = self.auth_response.slice();
if self.capability_flags.CLIENT_PLUGIN_AUTH_LENENC_CLIENT_DATA {
writer.write_length_encoded_string(auth_data)?;
} else if self.capability_flags.CLIENT_SECURE_CONNECTION {
writer.int1(u8::try_from(auth_data.len()).expect("int cast"))?;
writer.write(auth_data)?;
} else {
writer.write_z(auth_data)?;
}
if self.capability_flags.CLIENT_CONNECT_WITH_DB && self.database.slice().len() > 0 {
writer.write_z(self.database.slice())?;
}
if self.capability_flags.CLIENT_PLUGIN_AUTH {
writer.write_z(self.auth_plugin_name.slice())?;
}
if self.capability_flags.CLIENT_CONNECT_ATTRS {
let mut total_length: usize = 0;
for (key, value) in self.connect_attrs.iter() {
total_length += encode_length_int(key.len() as u64).len();
total_length += key.len();
total_length += encode_length_int(value.len() as u64).len();
total_length += value.len();
}
writer.write_length_encoded_int(total_length as u64)?;
for (key, value) in self.connect_attrs.iter() {
writer.write_length_encoded_string(key)?;
writer.write_length_encoded_string(value)?;
}
}
if self.capability_flags.CLIENT_ZSTD_COMPRESSION_ALGORITHM {
debug_assert!(false, "zstd compression algorithm is not supported");
}
packet.end()?;
Ok(())
}
pub fn write<Context: super::new_writer::WriterContext>(
&mut self,
writer: NewWriter<Context>,
) -> Result<(), bun_core::Error> {
self.write_internal(writer)
}
}