Skip to main content

servers/mysql/
server.rs

1// Copyright 2023 Greptime Team
2//
3// Licensed under the Apache License, Version 2.0 (the "License");
4// you may not use this file except in compliance with the License.
5// You may obtain a copy of the License at
6//
7//     http://www.apache.org/licenses/LICENSE-2.0
8//
9// Unless required by applicable law or agreed to in writing, software
10// distributed under the License is distributed on an "AS IS" BASIS,
11// WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied.
12// See the License for the specific language governing permissions and
13// limitations under the License.
14
15use 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
42// Default size of ResultSet write buffer: 100KB
43const 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
51/// [`MysqlSpawnRef`] stores arc refs
52/// that should be passed to new [`MysqlInstanceShim`]s.
53pub 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
77/// [`MysqlSpawnConfig`] stores config values
78/// which are used to initialize [`MysqlInstanceShim`]s.
79pub struct MysqlSpawnConfig {
80    // tls config
81    force_tls: bool,
82    tls: Arc<ReloadableTlsServerConfig>,
83    // keep-alive config
84    keep_alive_secs: u64,
85    // other shim config
86    reject_no_database: bool,
87    // prepared statement cache capacity
88    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    /// Enables ordinary-table batching for connections accepted by this server.
111    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"), // IoError doesn't impl ErrorExt.
171                    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                // This is a client-side error, we don't need to log it.
201            } else {
202                // TODO(LFC): Write this error to client as well, in MySQL text protocol.
203                // Looks like we have to expose opensrv-mysql's `PacketWriter`?
204                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}