1use 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
61enum Params<'a> {
63 ProtocolParams(Vec<ParamValue<'a>>),
65 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
78pub 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 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 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 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 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 fn plan(&self, stmt_key: &str) -> Option<SqlPlan> {
185 let guard = self.prepared_stmts.read();
186 guard.get(stmt_key).cloned()
187 }
188
189 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 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(¶m_types, param_num)?,
227 all_params_have_types(¶m_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, ¶ms, &timezone)
308 }
309 Params::CliParams(params) => {
310 replace_params_with_exprs(&plan, param_types, ¶ms, &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 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 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 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 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, ¶ms, &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 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 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 for (start, end, param) in replacements.into_iter().rev() {
828 query.replace_range(start..end, ¶m);
829 }
830
831 Ok(query)
832}
833
834fn location_to_byte_offset(query: &str, line: u64, column: u64) -> Option<usize> {
835 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 (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
964fn 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 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(¶m_types, 3).unwrap();
1085 assert_eq!(params.len(), 2);
1086 assert!(!all_params_have_types(¶m_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}