1use std::any::Any;
16use std::future::Future;
17use std::net::SocketAddr;
18use std::sync::Arc;
19
20use async_trait::async_trait;
21use auth::UserProviderRef;
22use catalog::process_manager::ProcessManagerRef;
23use common_runtime::Runtime;
24use common_runtime::runtime::RuntimeTrait;
25use common_telemetry::{debug, warn};
26use futures::StreamExt;
27use opensrv_mysql::{
28 AsyncMysqlIntermediary, IntermediaryOptions, plain_run_with_options, secure_run_with_options,
29};
30use snafu::ensure;
31use tokio;
32use tokio::io::BufWriter;
33use tokio::net::TcpStream;
34use tokio_rustls::rustls::ServerConfig;
35
36use crate::error::{Error, Result, TlsRequiredSnafu};
37use crate::mysql::handler::MysqlInstanceShim;
38use crate::query_handler::sql::ServerSqlQueryHandlerRef;
39use crate::server::{AbortableStream, BaseTcpServer, Server};
40use crate::tls::ReloadableTlsServerConfig;
41
42const DEFAULT_RESULT_SET_WRITE_BUFFER_SIZE: usize = 100 * 1024;
44
45const CLIENT_DISCONNECT_ERROR_KINDS: &[std::io::ErrorKind] = &[
46 std::io::ErrorKind::ConnectionAborted,
47 std::io::ErrorKind::ConnectionReset,
48 std::io::ErrorKind::BrokenPipe,
49];
50
51pub struct MysqlSpawnRef {
54 query_handler: ServerSqlQueryHandlerRef,
55 user_provider: Option<UserProviderRef>,
56}
57
58impl MysqlSpawnRef {
59 pub fn new(
60 query_handler: ServerSqlQueryHandlerRef,
61 user_provider: Option<UserProviderRef>,
62 ) -> MysqlSpawnRef {
63 MysqlSpawnRef {
64 query_handler,
65 user_provider,
66 }
67 }
68
69 fn query_handler(&self) -> ServerSqlQueryHandlerRef {
70 self.query_handler.clone()
71 }
72 fn user_provider(&self) -> Option<UserProviderRef> {
73 self.user_provider.clone()
74 }
75}
76
77pub struct MysqlSpawnConfig {
80 force_tls: bool,
82 tls: Arc<ReloadableTlsServerConfig>,
83 keep_alive_secs: u64,
85 reject_no_database: bool,
87 prepared_stmt_cache_size: usize,
89 batching_enabled: bool,
90}
91
92impl MysqlSpawnConfig {
93 pub fn new(
94 force_tls: bool,
95 tls: Arc<ReloadableTlsServerConfig>,
96 keep_alive_secs: u64,
97 reject_no_database: bool,
98 prepared_stmt_cache_size: usize,
99 ) -> MysqlSpawnConfig {
100 MysqlSpawnConfig {
101 force_tls,
102 tls,
103 keep_alive_secs,
104 reject_no_database,
105 prepared_stmt_cache_size,
106 batching_enabled: false,
107 }
108 }
109
110 pub fn with_batching_enabled(mut self, enabled: bool) -> Self {
112 self.batching_enabled = enabled;
113 self
114 }
115
116 fn tls(&self) -> Option<Arc<ServerConfig>> {
117 self.tls.get_config()
118 }
119}
120
121impl From<&MysqlSpawnConfig> for IntermediaryOptions {
122 fn from(value: &MysqlSpawnConfig) -> Self {
123 IntermediaryOptions {
124 reject_connection_on_dbname_absence: value.reject_no_database,
125 ..Default::default()
126 }
127 }
128}
129
130pub struct MysqlServer {
131 base_server: BaseTcpServer,
132 spawn_ref: Arc<MysqlSpawnRef>,
133 spawn_config: Arc<MysqlSpawnConfig>,
134 bind_addr: Option<SocketAddr>,
135 process_manager: Option<ProcessManagerRef>,
136}
137
138impl MysqlServer {
139 pub fn create_server(
140 io_runtime: Runtime,
141 spawn_ref: Arc<MysqlSpawnRef>,
142 spawn_config: Arc<MysqlSpawnConfig>,
143 process_manager: Option<ProcessManagerRef>,
144 ) -> Box<dyn Server> {
145 Box::new(MysqlServer {
146 base_server: BaseTcpServer::create_server("MySQL", io_runtime),
147 spawn_ref,
148 spawn_config,
149 bind_addr: None,
150 process_manager,
151 })
152 }
153
154 fn accept(
155 &self,
156 io_runtime: Runtime,
157 stream: AbortableStream,
158 process_manager: Option<ProcessManagerRef>,
159 ) -> impl Future<Output = ()> + use<> {
160 let spawn_ref = self.spawn_ref.clone();
161 let spawn_config = self.spawn_config.clone();
162
163 stream.for_each(move |tcp_stream| {
164 let spawn_ref = spawn_ref.clone();
165 let spawn_config = spawn_config.clone();
166 let io_runtime = io_runtime.clone();
167 let process_id = process_manager.as_ref().map(|p| p.next_id()).unwrap_or(8);
168 async move {
169 match tcp_stream {
170 Err(e) => warn!(e; "Broken pipe"), Ok(io_stream) => {
172 if let Err(e) = io_stream.set_nodelay(true) {
173 warn!(e; "Failed to set TCP nodelay");
174 }
175 io_runtime.spawn(async move {
176 if let Err(error) =
177 Self::handle(io_stream, spawn_ref, spawn_config, process_id).await
178 {
179 warn!(error; "Unexpected error when handling TcpStream");
180 };
181 });
182 }
183 };
184 }
185 })
186 }
187
188 async fn handle(
189 stream: TcpStream,
190 spawn_ref: Arc<MysqlSpawnRef>,
191 spawn_config: Arc<MysqlSpawnConfig>,
192 process_id: u32,
193 ) -> Result<()> {
194 debug!("MySQL connection coming from: {}", stream.peer_addr()?);
195 crate::metrics::METRIC_MYSQL_CONNECTIONS.inc();
196 if let Err(e) = Self::do_handle(stream, spawn_ref, spawn_config, process_id).await {
197 if let Error::InternalIo { error } = &e
198 && CLIENT_DISCONNECT_ERROR_KINDS.contains(&error.kind())
199 {
200 } else {
202 warn!(e; "Internal error occurred during query exec, server actively close the channel to let client try next time");
205 }
206 }
207 crate::metrics::METRIC_MYSQL_CONNECTIONS.dec();
208
209 Ok(())
210 }
211
212 async fn do_handle(
213 stream: TcpStream,
214 spawn_ref: Arc<MysqlSpawnRef>,
215 spawn_config: Arc<MysqlSpawnConfig>,
216 process_id: u32,
217 ) -> Result<()> {
218 let mut shim = MysqlInstanceShim::create(
219 spawn_ref.query_handler(),
220 spawn_ref.user_provider(),
221 stream.peer_addr()?,
222 process_id,
223 spawn_config.prepared_stmt_cache_size,
224 )
225 .with_batching_enabled(spawn_config.batching_enabled);
226 let (mut r, w) = stream.into_split();
227 let mut w = BufWriter::with_capacity(DEFAULT_RESULT_SET_WRITE_BUFFER_SIZE, w);
228
229 let ops = spawn_config.as_ref().into();
230
231 let (client_tls, init_params) =
232 AsyncMysqlIntermediary::init_before_ssl(&mut shim, &mut r, &mut w, &spawn_config.tls())
233 .await?;
234
235 ensure!(
236 !spawn_config.force_tls || client_tls,
237 TlsRequiredSnafu {
238 server: "mysql".to_owned()
239 }
240 );
241
242 match spawn_config.tls() {
243 Some(tls_conf) if client_tls => {
244 secure_run_with_options(shim, w, ops, tls_conf, init_params).await
245 }
246 _ => plain_run_with_options(shim, w, ops, init_params).await,
247 }
248 }
249}
250
251pub const MYSQL_SERVER: &str = "MYSQL_SERVER";
252
253#[async_trait]
254impl Server for MysqlServer {
255 async fn shutdown(&self) -> Result<()> {
256 self.base_server.shutdown().await
257 }
258
259 async fn start(&mut self, listening: SocketAddr) -> Result<()> {
260 let (stream, addr) = self
261 .base_server
262 .bind(listening, self.spawn_config.keep_alive_secs)
263 .await?;
264 let io_runtime = self.base_server.io_runtime();
265
266 let join_handle = common_runtime::spawn_global(self.accept(
267 io_runtime,
268 stream,
269 self.process_manager.clone(),
270 ));
271 self.base_server.start_with(join_handle).await?;
272
273 self.bind_addr = Some(addr);
274 Ok(())
275 }
276
277 fn name(&self) -> &str {
278 MYSQL_SERVER
279 }
280
281 fn bind_addr(&self) -> Option<SocketAddr> {
282 self.bind_addr
283 }
284
285 fn as_any(&self) -> &dyn Any {
286 self
287 }
288}
289
290#[cfg(test)]
291mod tests {
292 use super::CLIENT_DISCONNECT_ERROR_KINDS;
293
294 #[test]
295 fn test_client_disconnect_error_kinds() {
296 assert!(CLIENT_DISCONNECT_ERROR_KINDS.contains(&std::io::ErrorKind::ConnectionAborted));
297 assert!(CLIENT_DISCONNECT_ERROR_KINDS.contains(&std::io::ErrorKind::ConnectionReset));
298 assert!(CLIENT_DISCONNECT_ERROR_KINDS.contains(&std::io::ErrorKind::BrokenPipe));
299 }
300}