Skip to main content

servers/mysql/
handler.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::collections::HashMap;
16use std::net::SocketAddr;
17use std::sync::Arc;
18use std::sync::atomic::{AtomicU32, Ordering};
19use std::time::Duration;
20
21use ::auth::{BEARER_TOKEN_USER, Identity, MysqlAuthMethod, Password, UserProviderRef};
22use async_trait::async_trait;
23use chrono::{NaiveDate, NaiveDateTime};
24use common_catalog::parse_optional_catalog_and_schema_from_db_string;
25use common_error::ext::ErrorExt;
26use common_query::Output;
27use common_telemetry::{debug, error, tracing, warn};
28use common_time::Timezone;
29use datafusion_common::ParamValues;
30use datafusion_expr::LogicalPlan;
31use datatypes::prelude::ConcreteDataType;
32use datatypes::schema::Schema;
33use itertools::Itertools;
34use mysql_common::Value as MysqlValue;
35use opensrv_mysql::{
36    AsyncMysqlShim, Column, ErrorKind, InitWriter, ParamParser, ParamValue, QueryResultWriter,
37    StatementMetaWriter, ValueInner,
38};
39use parking_lot::RwLock;
40use query::planner::DfLogicalPlanner;
41use query::query_engine::DescribeResult;
42use rand::RngCore;
43use session::context::{Channel, QueryContextRef};
44use session::{Session, SessionRef};
45use snafu::{ResultExt, ensure};
46use sql::dialect::MySqlDialect;
47use sql::parser::{ParseOptions, ParserContext};
48use sql::statements::statement::Statement;
49use tokio::io::AsyncWrite;
50
51use crate::SqlPlan;
52use crate::error::{
53    self, DataFrameSnafu, InferParameterTypesSnafu, InvalidPrepareStatementSnafu, Result,
54};
55use crate::metrics::METRIC_AUTH_FAILURE;
56use crate::mysql::helper::{self, format_placeholder, transform_placeholders_with_count};
57use crate::mysql::writer;
58use crate::mysql::writer::{create_mysql_column, handle_err};
59use crate::query_handler::sql::ServerSqlQueryHandlerRef;
60
61/// Parameters for the prepared statement
62enum Params<'a> {
63    /// Parameters passed through protocol
64    ProtocolParams(Vec<ParamValue<'a>>),
65    /// Parameters passed through cli
66    CliParams(Vec<sql::ast::Expr>),
67}
68
69impl Params<'_> {
70    fn len(&self) -> usize {
71        match self {
72            Params::ProtocolParams(params) => params.len(),
73            Params::CliParams(params) => params.len(),
74        }
75    }
76}
77
78// An intermediate shim for executing MySQL queries.
79pub struct MysqlInstanceShim {
80    query_handler: ServerSqlQueryHandlerRef,
81    salt: [u8; 20],
82    session: SessionRef,
83    user_provider: Option<UserProviderRef>,
84    prepared_stmts: Arc<RwLock<HashMap<String, SqlPlan>>>,
85    prepared_stmts_counter: AtomicU32,
86    process_id: u32,
87    prepared_stmt_cache_size: usize,
88    batching_enabled: bool,
89}
90
91impl MysqlInstanceShim {
92    pub fn create(
93        query_handler: ServerSqlQueryHandlerRef,
94        user_provider: Option<UserProviderRef>,
95        client_addr: SocketAddr,
96        process_id: u32,
97        prepared_stmt_cache_size: usize,
98    ) -> MysqlInstanceShim {
99        // init a random salt
100        let mut bs = vec![0u8; 20];
101        let mut rng = rand::rng();
102        rng.fill_bytes(bs.as_mut());
103
104        let mut scramble: [u8; 20] = [0; 20];
105        for i in 0..20 {
106            scramble[i] = bs[i] & 0x7fu8;
107            if scramble[i] == b'\0' || scramble[i] == b'$' {
108                scramble[i] += 1;
109            }
110        }
111
112        MysqlInstanceShim {
113            query_handler,
114            salt: scramble,
115            session: Arc::new(Session::new(
116                Some(client_addr),
117                Channel::Mysql,
118                Default::default(),
119                process_id,
120            )),
121            user_provider,
122            prepared_stmts: Default::default(),
123            prepared_stmts_counter: AtomicU32::new(1),
124            process_id,
125            prepared_stmt_cache_size,
126            batching_enabled: false,
127        }
128    }
129
130    /// Enables ordinary-table batching for this connection.
131    pub fn with_batching_enabled(mut self, enabled: bool) -> Self {
132        self.batching_enabled = enabled;
133        self
134    }
135
136    fn new_query_context(&self) -> QueryContextRef {
137        let mut ctx = self.session.new_query_context();
138        Arc::make_mut(&mut ctx).set_batching_enabled(self.batching_enabled);
139        ctx
140    }
141
142    #[tracing::instrument(skip_all, name = "mysql::do_query")]
143    async fn do_query(&self, query: &str, query_ctx: QueryContextRef) -> Vec<Result<Output>> {
144        if let Some(output) =
145            crate::mysql::federated::check(query, query_ctx.clone(), self.session.clone())
146        {
147            vec![Ok(output)]
148        } else {
149            self.query_handler.do_query(query, query_ctx.clone()).await
150        }
151    }
152
153    /// Describe the statement
154    async fn do_describe(
155        &self,
156        statement: Statement,
157        query_ctx: QueryContextRef,
158    ) -> Result<Option<DescribeResult>> {
159        self.query_handler.do_describe(statement, query_ctx).await
160    }
161
162    /// Save query and logical plan with a given statement key
163    fn save_plan(&self, plan: SqlPlan, stmt_key: String) -> Result<()> {
164        let mut prepared_stmts = self.prepared_stmts.write();
165        let max_capacity = self.prepared_stmt_cache_size;
166
167        let is_update = prepared_stmts.contains_key(&stmt_key);
168
169        if !is_update && prepared_stmts.len() >= max_capacity {
170            return error::InternalSnafu {
171                err_msg: format!(
172                    "Prepared statement cache is full, max capacity: {}",
173                    max_capacity
174                ),
175            }
176            .fail();
177        }
178
179        let _ = prepared_stmts.insert(stmt_key, plan);
180        Ok(())
181    }
182
183    /// Retrieve the query and logical plan by a given statement key
184    fn plan(&self, stmt_key: &str) -> Option<SqlPlan> {
185        let guard = self.prepared_stmts.read();
186        guard.get(stmt_key).cloned()
187    }
188
189    /// Save the prepared statement and return the parameters and result columns
190    async fn do_prepare(
191        &mut self,
192        raw_query: &str,
193        query_ctx: QueryContextRef,
194        stmt_key: String,
195    ) -> Result<(Vec<Column>, Vec<Column>)> {
196        if crate::mysql::federated::check(raw_query, query_ctx.clone(), self.session.clone())
197            .is_some()
198        {
199            self.save_plan(SqlPlan::Shortcut(raw_query.to_string()), stmt_key)
200                .inspect_err(|e| {
201                    error!(e; "Failed to save prepared statement");
202                })?;
203            return Ok((vec![], vec![]));
204        }
205
206        let statement = validate_query(raw_query).await?;
207
208        // We have to transform the placeholder, because DataFusion only parses placeholders
209        // in the form of "$i", it can't process "?" right now.
210        let (statement, placeholder_count) = transform_placeholders_with_count(statement);
211        let param_num = placeholder_count + 1;
212
213        let describe_result = self
214            .do_describe(statement.clone(), query_ctx.clone())
215            .await?;
216        let plan = describe_result.map(|DescribeResult { logical_plan }| logical_plan);
217
218        let (params, can_cache_as_plan) = if let Some(plan) = &plan {
219            let param_types = DfLogicalPlanner::get_inferred_parameter_types(plan)
220                .context(InferParameterTypesSnafu)?
221                .into_iter()
222                .map(|(k, v)| (k, v.map(|v| ConcreteDataType::from_arrow_type(&v))))
223                .collect();
224
225            (
226                prepared_params(&param_types, param_num)?,
227                all_params_have_types(&param_types, param_num),
228            )
229        } else {
230            (dummy_params(param_num)?, false)
231        };
232
233        let columns =
234            plan.as_ref()
235                .map(|plan| {
236                    let schema: Schema = plan.schema().clone().try_into().map_err(
237                        |e: datatypes::error::Error| {
238                            error::InternalSnafu {
239                                err_msg: e.to_string(),
240                            }
241                            .build()
242                        },
243                    )?;
244                    schema
245                        .column_schemas()
246                        .iter()
247                        .map(|column_schema| {
248                            create_mysql_column(&column_schema.data_type, &column_schema.name)
249                        })
250                        .collect::<Result<Vec<_>>>()
251                })
252                .transpose()?
253                .unwrap_or_default();
254
255        match plan {
256            Some(plan) if can_cache_as_plan => {
257                self.save_plan(SqlPlan::Plan(plan, statement), stmt_key)
258                    .inspect_err(|e| {
259                        error!(e; "Failed to save prepared statement");
260                    })?;
261            }
262            _ => {
263                self.save_plan(
264                    SqlPlan::Statement(statement, raw_query.to_string()),
265                    stmt_key,
266                )
267                .inspect_err(|e| {
268                    error!(e; "Failed to save prepared statement");
269                })?;
270            }
271        }
272
273        Ok((params, columns))
274    }
275
276    async fn do_execute(
277        &mut self,
278        query_ctx: QueryContextRef,
279        stmt_key: String,
280        params: Params<'_>,
281    ) -> Result<Vec<std::result::Result<Output, error::Error>>> {
282        let sql_plan = match self.plan(&stmt_key) {
283            None => {
284                return error::PrepareStatementNotFoundSnafu { name: stmt_key }.fail();
285            }
286            Some(sql_plan) => sql_plan,
287        };
288
289        let outputs = match sql_plan {
290            SqlPlan::Plan(plan, stmt) => {
291                let param_types = DfLogicalPlanner::get_inferred_parameter_types(&plan)
292                    .context(InferParameterTypesSnafu)?
293                    .into_iter()
294                    .map(|(k, v)| (k, v.map(|v| ConcreteDataType::from_arrow_type(&v))))
295                    .collect::<HashMap<_, _>>();
296
297                if params.len() != param_types.len() {
298                    return error::InternalSnafu {
299                        err_msg: "Prepare statement params number mismatch".to_string(),
300                    }
301                    .fail();
302                }
303
304                let timezone = query_ctx.timezone();
305                let replaced_plan = match params {
306                    Params::ProtocolParams(params) => {
307                        replace_params_with_values(&plan, param_types, &params, &timezone)
308                    }
309                    Params::CliParams(params) => {
310                        replace_params_with_exprs(&plan, param_types, &params, &timezone)
311                    }
312                }?;
313
314                debug!(
315                    "Mysql execute prepared plan: {}",
316                    replaced_plan.display_indent()
317                );
318                vec![
319                    self.query_handler
320                        .do_exec_plan(replaced_plan, Some(stmt), query_ctx.clone())
321                        .await,
322                ]
323            }
324            SqlPlan::Shortcut(query) => {
325                if let Some(output) =
326                    crate::mysql::federated::check(&query, query_ctx.clone(), self.session.clone())
327                {
328                    vec![Ok(output)]
329                } else {
330                    self.do_query(&query, query_ctx.clone()).await
331                }
332            }
333            SqlPlan::Statement(stmt, query) => {
334                let param_strs = match params {
335                    Params::ProtocolParams(params) => {
336                        params.iter().map(convert_param_value_to_string).collect()
337                    }
338                    Params::CliParams(params) => params.iter().map(|x| x.to_string()).collect(),
339                };
340                debug!(
341                    "do_execute Replacing with Params: {:?}, Original Query: {}",
342                    param_strs, query
343                );
344                let query = replace_params(param_strs, stmt, query)?;
345                debug!("Mysql execute replaced query: {}", query);
346                self.do_query(&query, query_ctx.clone()).await
347            }
348            _ => {
349                return error::PrepareStatementNotFoundSnafu { name: stmt_key }.fail();
350            }
351        };
352
353        Ok(outputs)
354    }
355
356    /// Remove the prepared statement by a given statement key
357    fn do_close(&mut self, stmt_key: String) {
358        let mut guard = self.prepared_stmts.write();
359        let _ = guard.remove(&stmt_key);
360    }
361
362    fn auth_plugin(&self) -> &'static str {
363        self.user_provider
364            .as_ref()
365            .map(|provider| provider.mysql_auth_method())
366            .unwrap_or(MysqlAuthMethod::NativePassword)
367            .plugin_name()
368    }
369}
370
371#[async_trait]
372impl<W: AsyncWrite + Send + Sync + Unpin> AsyncMysqlShim<W> for MysqlInstanceShim {
373    type Error = error::Error;
374
375    fn version(&self) -> String {
376        std::env::var("GREPTIMEDB_MYSQL_SERVER_VERSION").unwrap_or_else(|_| "8.4.2".to_string())
377    }
378
379    fn connect_id(&self) -> u32 {
380        self.process_id
381    }
382
383    fn default_auth_plugin(&self) -> &str {
384        self.auth_plugin()
385    }
386
387    async fn auth_plugin_for_username(&self, user: &[u8]) -> &'static str {
388        if user == BEARER_TOKEN_USER.as_bytes() {
389            return MysqlAuthMethod::ClearPassword.plugin_name();
390        }
391        if let Some(provider) = &self.user_provider {
392            let username = String::from_utf8_lossy(user);
393            match provider.mysql_auth_method_for_user(&username).await {
394                Ok(method) => return method.plugin_name(),
395                // This hook cannot return an error. Keep the default challenge;
396                // authentication still validates the credentials separately.
397                Err(e) => warn!(e; "Failed to select MySQL authentication method"),
398            }
399        }
400        self.auth_plugin()
401    }
402
403    fn salt(&self) -> [u8; 20] {
404        self.salt
405    }
406
407    async fn authenticate(
408        &self,
409        auth_plugin: &str,
410        username: &[u8],
411        salt: &[u8],
412        auth_data: &[u8],
413    ) -> bool {
414        <Self as AsyncMysqlShim<W>>::authenticate_with_database(
415            self,
416            auth_plugin,
417            username,
418            salt,
419            auth_data,
420            None,
421        )
422        .await
423    }
424
425    async fn authenticate_with_database(
426        &self,
427        auth_plugin: &str,
428        username: &[u8],
429        salt: &[u8],
430        auth_data: &[u8],
431        database: Option<&[u8]>,
432    ) -> bool {
433        // if not specified then **greptime** will be used
434        let username = String::from_utf8_lossy(username);
435
436        let mut user_info = None;
437        let addr = self
438            .session
439            .conn_info()
440            .client_addr
441            .map(|addr| addr.to_string());
442        if let Some(user_provider) = &self.user_provider {
443            let result = if username.as_ref() == BEARER_TOKEN_USER {
444                if auth_plugin != MysqlAuthMethod::CLEAR_PASSWORD_PLUGIN {
445                    warn!("Bearer-token MySQL authentication requires mysql_clear_password");
446                    return false;
447                }
448                let token = auth_data.strip_suffix(&[0]).unwrap_or(auth_data);
449                let Ok(token) = std::str::from_utf8(token) else {
450                    warn!("Bearer token is not valid UTF-8");
451                    return false;
452                };
453                let catalog = if let Some(database) = database {
454                    let Ok(database) = std::str::from_utf8(database) else {
455                        warn!("MySQL database is not valid UTF-8");
456                        return false;
457                    };
458                    parse_optional_catalog_and_schema_from_db_string(database)
459                        .0
460                        .unwrap_or_else(|| self.session.catalog())
461                } else {
462                    self.session.catalog()
463                };
464                user_provider
465                    .authenticate_bearer_token(token, &catalog)
466                    .await
467            } else {
468                let user_id = Identity::UserId(&username, addr.as_deref());
469                let password = match auth_plugin {
470                    MysqlAuthMethod::NATIVE_PASSWORD_PLUGIN => {
471                        Password::MysqlNativePassword(auth_data, salt)
472                    }
473                    MysqlAuthMethod::CLEAR_PASSWORD_PLUGIN => {
474                        // The raw bytes received could be represented in C-like string, ended in '\0'.
475                        // We must "trim" it to get the real password string.
476                        let password = auth_data.strip_suffix(&[0]).unwrap_or(auth_data);
477                        Password::PlainText(String::from_utf8_lossy(password).to_string().into())
478                    }
479                    other => {
480                        error!("Unsupported mysql auth plugin: {}", other);
481                        return false;
482                    }
483                };
484                user_provider.authenticate(user_id, password).await
485            };
486            match result {
487                Ok(userinfo) => {
488                    user_info = Some(userinfo);
489                }
490                Err(e) => {
491                    METRIC_AUTH_FAILURE
492                        .with_label_values(&[e.status_code().as_ref()])
493                        .inc();
494                    warn!(e; "Failed to auth");
495                    return false;
496                }
497            };
498        }
499        let user_info =
500            user_info.unwrap_or_else(|| auth::userinfo_by_name(Some(username.to_string())));
501
502        self.session.set_user_info(user_info);
503
504        true
505    }
506
507    async fn on_prepare<'a>(
508        &'a mut self,
509        raw_query: &'a str,
510        w: StatementMetaWriter<'a, W>,
511    ) -> Result<()> {
512        let query_ctx = self.new_query_context();
513        let stmt_id = self.prepared_stmts_counter.fetch_add(1, Ordering::Relaxed);
514        let stmt_key = uuid::Uuid::from_u128(stmt_id as u128).to_string();
515        let (params, columns) = match self
516            .do_prepare(raw_query, query_ctx.clone(), stmt_key)
517            .await
518        {
519            Ok(x) => x,
520            Err(e) => {
521                let (kind, msg) = handle_err(e, query_ctx.clone());
522                w.error(kind, msg.as_bytes()).await?;
523                return Ok(());
524            }
525        };
526        debug!("on_prepare: Params: {:?}, Columns: {:?}", params, columns);
527        w.reply(stmt_id, &params, &columns).await?;
528        crate::metrics::METRIC_MYSQL_PREPARED_COUNT
529            .with_label_values(&[query_ctx.get_db_string().as_str()])
530            .inc();
531        return Ok(());
532    }
533
534    async fn on_execute<'a>(
535        &'a mut self,
536        stmt_id: u32,
537        p: ParamParser<'a>,
538        w: QueryResultWriter<'a, W>,
539    ) -> Result<()> {
540        self.session.clear_warnings();
541
542        let query_ctx = self.new_query_context();
543        let db = query_ctx.get_db_string();
544        let _timer = crate::metrics::METRIC_MYSQL_QUERY_TIMER
545            .with_label_values(&[crate::metrics::METRIC_MYSQL_BINQUERY, db.as_str()])
546            .start_timer();
547
548        let params: Vec<ParamValue> = p.into_iter().collect();
549        let stmt_key = uuid::Uuid::from_u128(stmt_id as u128).to_string();
550
551        let outputs = match self
552            .do_execute(query_ctx.clone(), stmt_key, Params::ProtocolParams(params))
553            .await
554        {
555            Ok(outputs) => outputs,
556            Err(e) => {
557                let (kind, err) = handle_err(e, query_ctx);
558                debug!(
559                    "Failed to execute prepared statement, kind: {:?}, err: {}",
560                    kind, err
561                );
562                w.error(kind, err.as_bytes()).await?;
563                return Ok(());
564            }
565        };
566
567        writer::write_output(w, query_ctx, self.session.clone(), outputs).await?;
568
569        Ok(())
570    }
571
572    async fn on_close<'a>(&'a mut self, stmt_id: u32)
573    where
574        W: 'async_trait,
575    {
576        let stmt_key = uuid::Uuid::from_u128(stmt_id as u128).to_string();
577        self.do_close(stmt_key);
578    }
579
580    #[tracing::instrument(skip_all, fields(protocol = "mysql"))]
581    async fn on_query<'a>(
582        &'a mut self,
583        query: &'a str,
584        writer: QueryResultWriter<'a, W>,
585    ) -> Result<()> {
586        let query_ctx = self.new_query_context();
587        let db = query_ctx.get_db_string();
588        let _timer = crate::metrics::METRIC_MYSQL_QUERY_TIMER
589            .with_label_values(&[crate::metrics::METRIC_MYSQL_TEXTQUERY, db.as_str()])
590            .start_timer();
591
592        // Clear warnings for non SHOW WARNINGS queries
593        let query_upcase = query.to_uppercase();
594        if !query_upcase.starts_with("SHOW WARNINGS") {
595            self.session.clear_warnings();
596        }
597
598        if query_upcase.starts_with("PREPARE ") {
599            match ParserContext::parse_mysql_prepare_stmt(query, query_ctx.sql_dialect()) {
600                Ok((stmt_name, stmt)) => {
601                    let prepare_results =
602                        self.do_prepare(&stmt, query_ctx.clone(), stmt_name).await;
603                    match prepare_results {
604                        Ok(_) => {
605                            let outputs = vec![Ok(Output::new_with_affected_rows(0))];
606                            writer::write_output(writer, query_ctx, self.session.clone(), outputs)
607                                .await?;
608                            return Ok(());
609                        }
610                        Err(e) => {
611                            writer
612                                .error(ErrorKind::ER_SP_BADSTATEMENT, e.output_msg().as_bytes())
613                                .await?;
614                            return Ok(());
615                        }
616                    }
617                }
618                Err(e) => {
619                    writer
620                        .error(ErrorKind::ER_PARSE_ERROR, e.output_msg().as_bytes())
621                        .await?;
622                    return Ok(());
623                }
624            }
625        } else if query_upcase.starts_with("EXECUTE ") {
626            match ParserContext::parse_mysql_execute_stmt(query, query_ctx.sql_dialect()) {
627                Ok((stmt_name, params)) => {
628                    let outputs = match self
629                        .do_execute(query_ctx.clone(), stmt_name, Params::CliParams(params))
630                        .await
631                    {
632                        Ok(outputs) => outputs,
633                        Err(e) => {
634                            let (kind, err) = handle_err(e, query_ctx);
635                            debug!(
636                                "Failed to execute prepared statement, kind: {:?}, err: {}",
637                                kind, err
638                            );
639                            writer.error(kind, err.as_bytes()).await?;
640                            return Ok(());
641                        }
642                    };
643                    writer::write_output(writer, query_ctx, self.session.clone(), outputs).await?;
644
645                    return Ok(());
646                }
647                Err(e) => {
648                    writer
649                        .error(ErrorKind::ER_PARSE_ERROR, e.output_msg().as_bytes())
650                        .await?;
651                    return Ok(());
652                }
653            }
654        } else if query_upcase.starts_with("DEALLOCATE ") {
655            match ParserContext::parse_mysql_deallocate_stmt(query, query_ctx.sql_dialect()) {
656                Ok(stmt_name) => {
657                    self.do_close(stmt_name);
658                    let outputs = vec![Ok(Output::new_with_affected_rows(0))];
659                    writer::write_output(writer, query_ctx, self.session.clone(), outputs).await?;
660                    return Ok(());
661                }
662                Err(e) => {
663                    writer
664                        .error(ErrorKind::ER_PARSE_ERROR, e.output_msg().as_bytes())
665                        .await?;
666                    return Ok(());
667                }
668            }
669        }
670
671        let outputs = self.do_query(query, query_ctx.clone()).await;
672        writer::write_output(writer, query_ctx, self.session.clone(), outputs).await?;
673
674        Ok(())
675    }
676
677    async fn on_init<'a>(&'a mut self, database: &'a str, w: InitWriter<'a, W>) -> Result<()> {
678        let (catalog_from_db, schema) = parse_optional_catalog_and_schema_from_db_string(database);
679        let catalog = if let Some(catalog) = &catalog_from_db {
680            catalog.clone()
681        } else {
682            self.session.catalog()
683        };
684
685        if !self
686            .query_handler
687            .is_valid_schema(&catalog, &schema)
688            .await?
689        {
690            return w
691                .error(
692                    ErrorKind::ER_WRONG_DB_NAME,
693                    format!("Unknown database '{}'", database).as_bytes(),
694                )
695                .await
696                .map_err(|e| e.into());
697        }
698
699        let user_info = &self.session.user_info();
700
701        if let Some(schema_validator) = &self.user_provider
702            && let Err(e) = schema_validator
703                .authorize(&catalog, &schema, user_info)
704                .await
705        {
706            METRIC_AUTH_FAILURE
707                .with_label_values(&[e.status_code().as_ref()])
708                .inc();
709            return w
710                .error(
711                    ErrorKind::ER_DBACCESS_DENIED_ERROR,
712                    e.output_msg().as_bytes(),
713                )
714                .await
715                .map_err(|e| e.into());
716        }
717
718        if catalog_from_db.is_some() {
719            self.session.set_catalog(catalog)
720        }
721        self.session.set_schema(schema);
722
723        w.ok().await.map_err(|e| e.into())
724    }
725}
726
727fn convert_param_value_to_string(param: &ParamValue) -> String {
728    match param.value.into_inner() {
729        ValueInner::Int(u) => u.to_string(),
730        ValueInner::UInt(u) => u.to_string(),
731        ValueInner::Double(u) => u.to_string(),
732        ValueInner::NULL => "NULL".to_string(),
733        // MySQL prepared fallback emits SQL text. Delegate bytes/string literal
734        // escaping to mysql_common. `false` means normal MySQL backslash escapes;
735        // if NO_BACKSLASH_ESCAPES is supported in this path later, wire the
736        // session SQL mode here.
737        ValueInner::Bytes(b) => MysqlValue::Bytes(b.to_vec()).as_sql(false),
738        ValueInner::Date(_) => format!("'{}'", NaiveDate::from(param.value)),
739        ValueInner::Datetime(_) => format!("'{}'", NaiveDateTime::from(param.value)),
740        ValueInner::Time(_) => format_duration(Duration::from(param.value)),
741    }
742}
743
744fn replace_params(params: Vec<String>, stmt: Statement, mut query: String) -> Result<String> {
745    let spans = helper::placeholder_spans(stmt);
746    ensure!(
747        spans.len() == params.len(),
748        error::InternalSnafu {
749            err_msg: format!(
750                "Prepared statement expected {} parameters but got {}",
751                spans.len(),
752                params.len()
753            )
754        }
755    );
756
757    let mut replacements = Vec::with_capacity(spans.len());
758    for span in spans {
759        let start = location_to_byte_offset(&query, span.start_line, span.start_column)
760            .ok_or_else(|| {
761                error::InternalSnafu {
762                    err_msg: format!(
763                        "Invalid placeholder start span: line {}, column {}",
764                        span.start_line, span.start_column
765                    ),
766                }
767                .build()
768            })?;
769        let end =
770            location_to_byte_offset(&query, span.end_line, span.end_column).ok_or_else(|| {
771                error::InternalSnafu {
772                    err_msg: format!(
773                        "Invalid placeholder end span: line {}, column {}",
774                        span.end_line, span.end_column
775                    ),
776                }
777                .build()
778            })?;
779        let param = span
780            .index
781            .checked_sub(1)
782            .and_then(|idx| params.get(idx))
783            .ok_or_else(|| {
784                error::InternalSnafu {
785                    err_msg: format!("Missing prepared statement parameter {}", span.index),
786                }
787                .build()
788            })?;
789
790        ensure!(
791            start < end && end <= query.len(),
792            error::InternalSnafu {
793                err_msg: format!(
794                    "Invalid placeholder byte span: {}..{} for query length {}",
795                    start,
796                    end,
797                    query.len()
798                )
799            }
800        );
801        ensure!(
802            query.get(start..end) == Some("?"),
803            error::InternalSnafu {
804                err_msg: format!(
805                    "Prepared statement placeholder span maps to {:?} instead of '?'",
806                    query.get(start..end)
807                )
808            }
809        );
810
811        replacements.push((start, end, param.clone()));
812    }
813
814    replacements.sort_unstable_by_key(|(start, _, _)| *start);
815    for windows in replacements.windows(2) {
816        ensure!(
817            windows[0].1 <= windows[1].0,
818            error::InternalSnafu {
819                err_msg: "Overlapping placeholder spans in prepared statement".to_string()
820            }
821        );
822    }
823
824    // All spans are computed against the original query. Apply replacements
825    // from right to left so changing one parameter's string length never shifts
826    // the byte offsets of placeholders that have not been replaced yet.
827    for (start, end, param) in replacements.into_iter().rev() {
828        query.replace_range(start..end, &param);
829    }
830
831    Ok(query)
832}
833
834fn location_to_byte_offset(query: &str, line: u64, column: u64) -> Option<usize> {
835    // sqlparser spans are 1-based line/column locations, and columns advance by
836    // Rust `char`s rather than bytes. Convert them to byte offsets before using
837    // `String::replace_range` on the original SQL text.
838    if line == 0 || column == 0 {
839        return None;
840    }
841
842    let mut current_line = 1;
843    let mut current_column = 1;
844    for (index, ch) in query.char_indices() {
845        if current_line == line && current_column == column {
846            return Some(index);
847        }
848
849        if ch == '\n' {
850            current_line += 1;
851            current_column = 1;
852        } else {
853            current_column += 1;
854        }
855    }
856
857    // The exclusive end location of a trailing placeholder points just past
858    // the last character, for example the end span of `SELECT ?`.
859    (current_line == line && current_column == column).then_some(query.len())
860}
861
862fn format_duration(duration: Duration) -> String {
863    let seconds = duration.as_secs() % 60;
864    let minutes = (duration.as_secs() / 60) % 60;
865    let hours = (duration.as_secs() / 60) / 60;
866    format!("'{}:{}:{}'", hours, minutes, seconds)
867}
868
869fn replace_params_with_values(
870    plan: &LogicalPlan,
871    param_types: HashMap<String, Option<ConcreteDataType>>,
872    params: &[ParamValue],
873    timezone: &Timezone,
874) -> Result<LogicalPlan> {
875    debug_assert_eq!(param_types.len(), params.len());
876
877    debug!(
878        "replace_params_with_values(param_types: {:#?}, params: {:#?}, plan: {:#?})",
879        param_types,
880        params
881            .iter()
882            .map(|x| format!("({:?}, {:?})", x.value, x.coltype))
883            .join(", "),
884        plan
885    );
886
887    let mut values = Vec::with_capacity(params.len());
888
889    for (i, param) in params.iter().enumerate() {
890        if let Some(Some(t)) = param_types.get(&format_placeholder(i + 1)) {
891            let value = helper::convert_value(param, t, timezone)?;
892
893            values.push(value.into());
894        }
895    }
896
897    plan.clone()
898        .replace_params_with_values(&ParamValues::List(values.clone()))
899        .context(DataFrameSnafu)
900}
901
902fn replace_params_with_exprs(
903    plan: &LogicalPlan,
904    param_types: HashMap<String, Option<ConcreteDataType>>,
905    params: &[sql::ast::Expr],
906    timezone: &Timezone,
907) -> Result<LogicalPlan> {
908    debug_assert_eq!(param_types.len(), params.len());
909
910    debug!(
911        "replace_params_with_exprs(param_types: {:#?}, params: {:#?}, plan: {:#?})",
912        param_types,
913        params.iter().map(|x| format!("({:?})", x)).join(", "),
914        plan
915    );
916
917    let mut values = Vec::with_capacity(params.len());
918
919    for (i, param) in params.iter().enumerate() {
920        if let Some(Some(t)) = param_types.get(&format_placeholder(i + 1)) {
921            let value = helper::convert_expr_to_scalar_value(param, t, timezone)?;
922
923            values.push(value.into());
924        }
925    }
926
927    plan.clone()
928        .replace_params_with_values(&ParamValues::List(values.clone()))
929        .context(DataFrameSnafu)
930}
931
932async fn validate_query(query: &str) -> Result<Statement> {
933    let statement =
934        ParserContext::create_with_dialect(query, &MySqlDialect {}, ParseOptions::default());
935    let mut statement = statement.map_err(|e| {
936        InvalidPrepareStatementSnafu {
937            err_msg: e.output_msg(),
938        }
939        .build()
940    })?;
941
942    ensure!(
943        statement.len() == 1,
944        InvalidPrepareStatementSnafu {
945            err_msg: "prepare statement only support single statement".to_string(),
946        }
947    );
948
949    let statement = statement.remove(0);
950
951    Ok(statement)
952}
953
954fn dummy_params(index: usize) -> Result<Vec<Column>> {
955    let mut params = Vec::with_capacity(index - 1);
956
957    for _ in 1..index {
958        params.push(create_mysql_column(&ConcreteDataType::null_datatype(), "")?);
959    }
960
961    Ok(params)
962}
963
964/// Parameters that the client must provide when executing the prepared statement.
965fn prepared_params(
966    param_types: &HashMap<String, Option<ConcreteDataType>>,
967    param_num: usize,
968) -> Result<Vec<Column>> {
969    let mut params = Vec::with_capacity(param_num - 1);
970
971    // Placeholder index starts from 1
972    for i in 1..param_num {
973        let column = if let Some(Some(t)) = param_types.get(&format_placeholder(i)) {
974            create_mysql_column(t, "")?
975        } else {
976            create_mysql_column(&ConcreteDataType::null_datatype(), "")?
977        };
978        params.push(column);
979    }
980
981    Ok(params)
982}
983
984fn all_params_have_types(
985    param_types: &HashMap<String, Option<ConcreteDataType>>,
986    param_num: usize,
987) -> bool {
988    param_types.len() == param_num - 1
989        && (1..param_num).all(|i| matches!(param_types.get(&format_placeholder(i)), Some(Some(_))))
990}
991
992#[cfg(test)]
993mod tests {
994    use std::sync::Arc;
995
996    use async_trait::async_trait;
997    use common_query::Output;
998    use datafusion_expr::LogicalPlan;
999    use query::parser::PromQuery;
1000    use query::query_engine::DescribeResult;
1001    use session::context::QueryContext;
1002    use sql::statements::statement::Statement;
1003
1004    use super::*;
1005    use crate::error::Result;
1006    use crate::query_handler::sql::SqlQueryHandler;
1007
1008    struct DummyQueryHandler;
1009
1010    #[async_trait]
1011    impl SqlQueryHandler for DummyQueryHandler {
1012        async fn do_query(&self, _: &str, _: QueryContextRef) -> Vec<Result<Output>> {
1013            unimplemented!()
1014        }
1015
1016        async fn do_analyze_stream_query(&self, _: &str, _: QueryContextRef) -> Result<Output> {
1017            unimplemented!()
1018        }
1019
1020        async fn do_promql_query(&self, _: &PromQuery, _: QueryContextRef) -> Vec<Result<Output>> {
1021            unimplemented!()
1022        }
1023
1024        async fn do_exec_plan(
1025            &self,
1026            _: LogicalPlan,
1027            _: Option<Statement>,
1028            _: QueryContextRef,
1029        ) -> Result<Output> {
1030            unimplemented!()
1031        }
1032
1033        async fn do_describe(
1034            &self,
1035            _: Statement,
1036            _: QueryContextRef,
1037        ) -> Result<Option<DescribeResult>> {
1038            unimplemented!()
1039        }
1040
1041        async fn is_valid_schema(&self, _: &str, _: &str) -> Result<bool> {
1042            Ok(true)
1043        }
1044    }
1045
1046    fn create_shim() -> MysqlInstanceShim {
1047        MysqlInstanceShim::create(
1048            Arc::new(DummyQueryHandler),
1049            None,
1050            "127.0.0.1:3306".parse().unwrap(),
1051            1,
1052            1024,
1053        )
1054    }
1055
1056    #[test]
1057    fn test_batching_context() {
1058        for enabled in [false, true] {
1059            let shim = create_shim().with_batching_enabled(enabled);
1060            let ctx = shim.new_query_context();
1061            assert_eq!(ctx.batching_enabled(), enabled);
1062            assert!(!ctx.logical_batching_enabled());
1063            assert_eq!(ctx.channel(), Channel::Mysql);
1064        }
1065    }
1066
1067    fn statement_with_transformed_placeholders(query: &str) -> Statement {
1068        let mut statements =
1069            ParserContext::create_with_dialect(query, &MySqlDialect {}, ParseOptions::default())
1070                .unwrap();
1071        assert_eq!(statements.len(), 1);
1072        transform_placeholders_with_count(statements.remove(0)).0
1073    }
1074
1075    #[test]
1076    fn test_prepared_params_keep_unknown_type_placeholders() {
1077        let mut param_types = HashMap::new();
1078        param_types.insert(format_placeholder(1), None);
1079        param_types.insert(
1080            format_placeholder(2),
1081            Some(ConcreteDataType::int32_datatype()),
1082        );
1083
1084        let params = prepared_params(&param_types, 3).unwrap();
1085        assert_eq!(params.len(), 2);
1086        assert!(!all_params_have_types(&param_types, 3));
1087    }
1088
1089    #[test]
1090    fn test_replace_params_by_placeholder_span() {
1091        let query = "SELECT ?, ?".to_string();
1092        let stmt = statement_with_transformed_placeholders(&query);
1093        let params = vec!["'$2 should stay'".to_string(), "'value'".to_string()];
1094
1095        assert_eq!(
1096            "SELECT '$2 should stay', 'value'",
1097            replace_params(params, stmt, query).unwrap()
1098        );
1099
1100        let query = "SELECT ?, ?, ?".to_string();
1101        let stmt = statement_with_transformed_placeholders(&query);
1102        let params = vec![
1103            "'much longer than a placeholder'".to_string(),
1104            "0".to_string(),
1105            "'also much longer than a placeholder'".to_string(),
1106        ];
1107
1108        assert_eq!(
1109            "SELECT 'much longer than a placeholder', 0, 'also much longer than a placeholder'",
1110            replace_params(params, stmt, query).unwrap()
1111        );
1112
1113        let query = "SELECT '$1', \"$2\", `$3`, ?, ?".to_string();
1114        let stmt = statement_with_transformed_placeholders(&query);
1115        let params = vec!["'1'".to_string(), "'2'".to_string()];
1116
1117        assert_eq!(
1118            "SELECT '$1', \"$2\", `$3`, '1', '2'",
1119            replace_params(params, stmt, query).unwrap()
1120        );
1121
1122        let query = "SELECT /* ? */ ? -- ?\n, ?".to_string();
1123        let stmt = statement_with_transformed_placeholders(&query);
1124        let params = vec!["'first'".to_string(), "'second'".to_string()];
1125
1126        assert_eq!(
1127            "SELECT /* ? */ 'first' -- ?\n, 'second'",
1128            replace_params(params, stmt, query).unwrap()
1129        );
1130
1131        let query = "SELECT '中文', ?".to_string();
1132        let stmt = statement_with_transformed_placeholders(&query);
1133        let params = vec!["'value'".to_string()];
1134
1135        assert_eq!(
1136            "SELECT '中文', 'value'",
1137            replace_params(params, stmt, query).unwrap()
1138        );
1139
1140        let query = "SELECT '中文',\n  ?".to_string();
1141        let stmt = statement_with_transformed_placeholders(&query);
1142        let params = vec!["'value'".to_string()];
1143
1144        assert_eq!(
1145            "SELECT '中文',\n  'value'",
1146            replace_params(params, stmt, query).unwrap()
1147        );
1148
1149        let query = "SELECT 'x'\r\n, ?".to_string();
1150        let stmt = statement_with_transformed_placeholders(&query);
1151        let params = vec!["'crlf'".to_string()];
1152
1153        assert_eq!(
1154            "SELECT 'x'\r\n, 'crlf'",
1155            replace_params(params, stmt, query).unwrap()
1156        );
1157
1158        let query = "SELECT\t?".to_string();
1159        let stmt = statement_with_transformed_placeholders(&query);
1160        let params = vec!["NULL".to_string()];
1161
1162        assert_eq!("SELECT\tNULL", replace_params(params, stmt, query).unwrap());
1163
1164        let query = "SELECT CAST(? AS INT64), ? + (SELECT ?)".to_string();
1165        let stmt = statement_with_transformed_placeholders(&query);
1166        let params = vec!["1".to_string(), "2".to_string(), "3".to_string()];
1167
1168        assert_eq!(
1169            "SELECT CAST(1 AS INT64), 2 + (SELECT 3)",
1170            replace_params(params, stmt, query).unwrap()
1171        );
1172
1173        let query = "SET time_zone = ?".to_string();
1174        let stmt = statement_with_transformed_placeholders(&query);
1175        let params = vec!["'UTC'".to_string()];
1176
1177        assert_eq!(
1178            "SET time_zone = 'UTC'",
1179            replace_params(params, stmt, query).unwrap()
1180        );
1181    }
1182
1183    #[tokio::test]
1184    async fn test_prepare_federated_query() {
1185        let mut shim = create_shim();
1186        let query_ctx = QueryContext::arc();
1187        let stmt_key = "test_federated".to_string();
1188
1189        let (params, columns) = shim
1190            .do_prepare(
1191                "SELECT @@version_comment",
1192                query_ctx.clone(),
1193                stmt_key.clone(),
1194            )
1195            .await
1196            .unwrap();
1197
1198        assert!(params.is_empty());
1199        assert!(columns.is_empty());
1200
1201        let plan = shim.plan(&stmt_key).unwrap();
1202        assert!(matches!(plan, SqlPlan::Shortcut(q) if q == "SELECT @@version_comment"));
1203    }
1204
1205    #[tokio::test]
1206    async fn test_execute_federated_shortcut() {
1207        let mut shim = create_shim();
1208        let query_ctx = QueryContext::arc();
1209        let stmt_key = "test_federated_exec".to_string();
1210
1211        shim.do_prepare(
1212            "SELECT @@version_comment",
1213            query_ctx.clone(),
1214            stmt_key.clone(),
1215        )
1216        .await
1217        .unwrap();
1218
1219        let outputs = shim
1220            .do_execute(query_ctx.clone(), stmt_key, Params::CliParams(vec![]))
1221            .await
1222            .unwrap();
1223
1224        assert_eq!(outputs.len(), 1);
1225        let output = outputs.into_iter().next().unwrap().unwrap();
1226        let pretty = output.data.pretty_print().await;
1227        assert!(pretty.contains("GreptimeDB"));
1228    }
1229
1230    #[tokio::test]
1231    async fn test_prepare_non_federated_query_not_shortcut() {
1232        let mut shim = create_shim();
1233        let query_ctx = QueryContext::arc();
1234        let stmt_key = "test_non_federated".to_string();
1235
1236        let result = shim
1237            .do_prepare("SET NAMES utf8", query_ctx.clone(), stmt_key.clone())
1238            .await;
1239
1240        assert!(result.is_ok());
1241        let plan = shim.plan(&stmt_key).unwrap();
1242        assert!(matches!(plan, SqlPlan::Shortcut(_)));
1243    }
1244
1245    #[tokio::test]
1246    async fn test_execute_set_shortcut() {
1247        let mut shim = create_shim();
1248        let query_ctx = QueryContext::arc();
1249        let stmt_key = "test_set_shortcut".to_string();
1250
1251        shim.do_prepare("SET NAMES utf8", query_ctx.clone(), stmt_key.clone())
1252            .await
1253            .unwrap();
1254
1255        let outputs = shim
1256            .do_execute(query_ctx.clone(), stmt_key, Params::CliParams(vec![]))
1257            .await
1258            .unwrap();
1259
1260        assert_eq!(outputs.len(), 1);
1261        let output = outputs.into_iter().next().unwrap().unwrap();
1262        match output.data {
1263            common_query::OutputData::RecordBatches(batches) => {
1264                let total_rows: usize = batches.iter().map(|b| b.num_rows()).sum();
1265                assert_eq!(total_rows, 0);
1266            }
1267            other => panic!("Expected RecordBatches, got {:?}", other),
1268        }
1269    }
1270}