Skip to main content

session/
lib.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
15pub mod context;
16pub mod hints;
17pub mod protocol_ctx;
18pub mod query_id;
19pub mod session_config;
20pub mod table_name;
21
22use std::collections::{HashMap, VecDeque};
23use std::net::SocketAddr;
24use std::sync::{Arc, RwLock};
25use std::time::Duration;
26
27use auth::UserInfoRef;
28use common_catalog::build_db_string;
29use common_catalog::consts::{DEFAULT_CATALOG_NAME, DEFAULT_SCHEMA_NAME};
30use common_recordbatch::cursor::RecordBatchStreamCursor;
31pub use common_session::ReadPreference;
32use common_time::Timezone;
33use common_time::timezone::get_timezone;
34use context::{ConfigurationVariables, QueryContextBuilder};
35use derive_more::Debug;
36
37use crate::context::{Channel, ConnInfo, QueryContextRef, dialect_for_channel};
38
39/// Maximum number of warnings to store per session (similar to MySQL's max_error_count)
40const MAX_WARNINGS: usize = 64;
41
42/// Session for persistent connection such as MySQL, PostgreSQL etc.
43#[derive(Debug)]
44pub struct Session {
45    catalog: RwLock<String>,
46    mutable_inner: Arc<RwLock<MutableInner>>,
47    conn_info: ConnInfo,
48    configuration_variables: Arc<ConfigurationVariables>,
49    // the process id to use when killing the query
50    process_id: u32,
51}
52
53pub type SessionRef = Arc<Session>;
54
55/// A container for mutable items in query context
56#[derive(Debug, Clone)]
57pub(crate) struct MutableInner {
58    schema: String,
59    user_info: UserInfoRef,
60    timezone: Timezone,
61    query_timeout: Option<Duration>,
62    read_preference: ReadPreference,
63    /// Request-level WAL policy for ordinary inserts.
64    skip_wal: bool,
65    #[debug(skip)]
66    pub(crate) cursors: HashMap<String, Arc<RecordBatchStreamCursor>>,
67    /// Warning messages for MySQL SHOW WARNINGS support
68    warnings: VecDeque<String>,
69}
70
71impl Default for MutableInner {
72    fn default() -> Self {
73        Self {
74            schema: DEFAULT_SCHEMA_NAME.into(),
75            user_info: auth::userinfo_by_name(None),
76            timezone: get_timezone(None).clone(),
77            query_timeout: None,
78            read_preference: ReadPreference::Leader,
79            skip_wal: false,
80            cursors: HashMap::with_capacity(0),
81            warnings: VecDeque::new(),
82        }
83    }
84}
85
86impl Session {
87    pub fn new(
88        addr: Option<SocketAddr>,
89        channel: Channel,
90        configuration_variables: ConfigurationVariables,
91        process_id: u32,
92    ) -> Self {
93        Session {
94            catalog: RwLock::new(DEFAULT_CATALOG_NAME.into()),
95            conn_info: ConnInfo::new(addr, channel),
96            configuration_variables: Arc::new(configuration_variables),
97            mutable_inner: Arc::new(RwLock::new(MutableInner::default())),
98            process_id,
99        }
100    }
101
102    pub fn new_query_context(&self) -> QueryContextRef {
103        QueryContextBuilder::default()
104            // catalog is not allowed for update in query context so we use
105            // string here
106            .current_catalog(self.catalog.read().unwrap().clone())
107            .mutable_session_data(self.mutable_inner.clone())
108            .sql_dialect(dialect_for_channel(self.conn_info.channel))
109            .configuration_parameter(self.configuration_variables.clone())
110            .channel(self.conn_info.channel)
111            .process_id(self.process_id)
112            .conn_info(self.conn_info.clone())
113            .build()
114            .into()
115    }
116
117    /// Cursors are shared across query contexts created from this session.
118    pub fn get_cursor(&self, name: &str) -> Option<Arc<RecordBatchStreamCursor>> {
119        let guard = self.mutable_inner.read().unwrap();
120        guard.cursors.get(name).cloned()
121    }
122
123    pub fn conn_info(&self) -> &ConnInfo {
124        &self.conn_info
125    }
126
127    pub fn timezone(&self) -> Timezone {
128        self.mutable_inner.read().unwrap().timezone.clone()
129    }
130
131    pub fn read_preference(&self) -> ReadPreference {
132        self.mutable_inner.read().unwrap().read_preference
133    }
134
135    pub fn set_timezone(&self, tz: Timezone) {
136        let mut inner = self.mutable_inner.write().unwrap();
137        inner.timezone = tz;
138    }
139
140    pub fn set_read_preference(&self, read_preference: ReadPreference) {
141        self.mutable_inner.write().unwrap().read_preference = read_preference;
142    }
143
144    pub fn user_info(&self) -> UserInfoRef {
145        self.mutable_inner.read().unwrap().user_info.clone()
146    }
147
148    pub fn set_user_info(&self, user_info: UserInfoRef) {
149        self.mutable_inner.write().unwrap().user_info = user_info;
150    }
151
152    pub fn set_catalog(&self, catalog: String) {
153        *self.catalog.write().unwrap() = catalog;
154    }
155
156    pub fn catalog(&self) -> String {
157        self.catalog.read().unwrap().clone()
158    }
159
160    pub fn schema(&self) -> String {
161        self.mutable_inner.read().unwrap().schema.clone()
162    }
163
164    pub fn set_schema(&self, schema: String) {
165        self.mutable_inner.write().unwrap().schema = schema;
166    }
167
168    pub fn get_db_string(&self) -> String {
169        build_db_string(&self.catalog(), &self.schema())
170    }
171
172    pub fn process_id(&self) -> u32 {
173        self.process_id
174    }
175
176    pub fn warnings_count(&self) -> usize {
177        self.mutable_inner.read().unwrap().warnings.len()
178    }
179
180    pub fn warnings(&self) -> Vec<String> {
181        self.mutable_inner
182            .read()
183            .unwrap()
184            .warnings
185            .iter()
186            .cloned()
187            .collect()
188    }
189
190    /// Add a warning message. If the limit is reached, discard the oldest warning.
191    pub fn add_warning(&self, warning: String) {
192        let mut inner = self.mutable_inner.write().unwrap();
193        if inner.warnings.len() >= MAX_WARNINGS {
194            inner.warnings.pop_front();
195        }
196        inner.warnings.push_back(warning);
197    }
198
199    pub fn clear_warnings(&self) {
200        let mut inner = self.mutable_inner.write().unwrap();
201        if inner.warnings.is_empty() {
202            return;
203        }
204        inner.warnings.clear();
205    }
206}