Skip to main content

risingwave_frontend/handler/
alter_source_column.rs

1// Copyright 2023 RisingWave Labs
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 pgwire::pg_response::{PgResponse, StatementType};
16use risingwave_common::catalog::max_column_id;
17use risingwave_connector::source::{SourceEncode, SourceStruct, extract_source_struct};
18use risingwave_sqlparser::ast::{AlterSourceOperation, ObjectName};
19
20use super::create_source::{generate_stream_graph_for_source, reject_variant_columns};
21use super::create_table::bind_sql_columns;
22use super::{HandlerArgs, RwPgResponse};
23use crate::Binder;
24use crate::catalog::root_catalog::SchemaPath;
25use crate::error::{ErrorCode, Result, RwError};
26
27// Note for future drop column:
28// 1. Dependencies of generated columns
29
30/// Handle `ALTER TABLE [ADD] COLUMN` statements.
31pub async fn handle_alter_source_column(
32    handler_args: HandlerArgs,
33    source_name: ObjectName,
34    operation: AlterSourceOperation,
35) -> Result<RwPgResponse> {
36    // Get original definition
37    let session = handler_args.session.clone();
38    let db_name = &session.database();
39    let (schema_name, real_source_name) =
40        Binder::resolve_schema_qualified_name(db_name, &source_name)?;
41    let search_path = session.config().search_path();
42    let user_name = &session.user_name();
43
44    let schema_path = SchemaPath::new(schema_name.as_deref(), &search_path, user_name);
45
46    let mut catalog = {
47        let reader = session.env().catalog_reader().read_guard();
48        let (source, schema_name) =
49            reader.get_source_by_name(db_name, schema_path, &real_source_name)?;
50        session.check_privilege_for_drop_alter(schema_name, &**source)?;
51
52        (**source).clone()
53    };
54
55    if catalog.associated_table_id.is_some() {
56        return Err(ErrorCode::NotSupported(
57            "alter table with connector with ALTER SOURCE statement".to_owned(),
58            "try to use ALTER TABLE instead".to_owned(),
59        )
60        .into());
61    };
62    if catalog.is_cdc_table_source() {
63        return Err(ErrorCode::NotSupported(
64            "altering columns of a CDC table source is not supported".to_owned(),
65            "Drop and recreate the CDC table source".to_owned(),
66        )
67        .into());
68    }
69
70    // Currently only allow source without schema registry
71    let SourceStruct { encode, .. } = extract_source_struct(&catalog.info)?;
72    match encode {
73        SourceEncode::Avro | SourceEncode::Protobuf => {
74            return Err(ErrorCode::NotSupported(
75                "alter source with schema registry".to_owned(),
76                "try `ALTER SOURCE .. FORMAT .. ENCODE .. (...)` instead".to_owned(),
77            )
78            .into());
79        }
80        SourceEncode::Json if catalog.info.use_schema_registry => {
81            return Err(ErrorCode::NotSupported(
82                "alter source with schema registry".to_owned(),
83                "try `ALTER SOURCE .. FORMAT .. ENCODE .. (...)` instead".to_owned(),
84            )
85            .into());
86        }
87        SourceEncode::Invalid | SourceEncode::Native | SourceEncode::None => {
88            return Err(RwError::from(ErrorCode::NotSupported(
89                format!("alter source with encode {:?}", encode),
90                "Only source with encode JSON | BYTES | CSV | PARQUET can be altered".into(),
91            )));
92        }
93        SourceEncode::Json | SourceEncode::Csv | SourceEncode::Bytes | SourceEncode::Parquet => {}
94    }
95
96    let columns = &mut catalog.columns;
97    match operation {
98        AlterSourceOperation::AddColumn { column_def } => {
99            let new_column_name = column_def.name.real_value();
100            if columns
101                .iter()
102                .any(|c| c.column_desc.name == new_column_name)
103            {
104                Err(ErrorCode::InvalidInputSyntax(format!(
105                    "column \"{new_column_name}\" of source \"{source_name}\" already exists"
106                )))?
107            }
108
109            // add column name is from user, so we still have check for reserved column name
110            let mut bound_column = bind_sql_columns(&[column_def], false)?.remove(0);
111            // PARQUET reads variant via the extension type; other alterable encodings cannot
112            // produce variant values, and this path bypasses the CREATE-time gate.
113            if encode != SourceEncode::Parquet {
114                reject_variant_columns(
115                    std::slice::from_ref(&bound_column),
116                    "for this source encoding",
117                )?;
118            }
119            bound_column.column_desc.column_id = max_column_id(columns).next();
120            columns.push(bound_column);
121            // No need to update the definition here. It will be done by purification later.
122        }
123        _ => unreachable!(),
124    }
125
126    // update version
127    catalog.version += 1;
128    catalog.fill_purified_create_sql();
129
130    let catalog_writer = session.catalog_writer()?;
131    if catalog.info.is_shared() {
132        let graph = generate_stream_graph_for_source(handler_args, catalog.clone())?;
133        catalog_writer
134            .replace_source(catalog.to_prost(), graph)
135            .await?
136    } else {
137        catalog_writer.alter_source(catalog.to_prost()).await?
138    };
139
140    Ok(PgResponse::empty_result(StatementType::ALTER_SOURCE))
141}
142
143#[cfg(test)]
144pub mod tests {
145    use std::collections::BTreeMap;
146
147    use risingwave_common::catalog::{DEFAULT_DATABASE_NAME, DEFAULT_SCHEMA_NAME};
148
149    use crate::catalog::root_catalog::SchemaPath;
150    use crate::test_utils::LocalFrontend;
151
152    #[tokio::test]
153    async fn test_alter_source_column_handler() {
154        let frontend = LocalFrontend::new(Default::default()).await;
155        let session = frontend.session_ref();
156        let schema_path = SchemaPath::Name(DEFAULT_SCHEMA_NAME);
157
158        let sql = r#"create source s_shared (v1 int) with (
159            connector = 'kafka',
160            topic = 'abc',
161            properties.bootstrap.server = 'localhost:29092',
162        ) FORMAT PLAIN ENCODE JSON;"#;
163
164        frontend
165            .run_sql_with_session(session.clone(), sql)
166            .await
167            .unwrap();
168
169        frontend
170            .run_sql_with_session(session.clone(), "SET streaming_use_shared_source TO false;")
171            .await
172            .unwrap();
173        let sql = r#"create source s (v1 int) with (
174            connector = 'kafka',
175            topic = 'abc',
176            properties.bootstrap.server = 'localhost:29092',
177          ) FORMAT PLAIN ENCODE JSON;"#;
178
179        frontend
180            .run_sql_with_session(session.clone(), sql)
181            .await
182            .unwrap();
183
184        let get_source = |name: &str| {
185            let catalog_reader = session.env().catalog_reader().read_guard();
186            catalog_reader
187                .get_source_by_name(DEFAULT_DATABASE_NAME, schema_path, name)
188                .unwrap()
189                .0
190                .clone()
191        };
192
193        let source = get_source("s");
194
195        let sql = "alter source s_shared add column v2 varchar;";
196        frontend.run_sql(sql).await.unwrap();
197
198        let altered_source = get_source("s_shared");
199        let altered_columns: BTreeMap<_, _> = altered_source
200            .columns
201            .iter()
202            .map(|col| (col.name(), (col.data_type().clone(), col.column_id())))
203            .collect();
204
205        // Check the new column is added.
206        // Check the old columns and IDs are not changed.
207        expect_test::expect![[r#"
208            {
209                "_row_id": (
210                    Serial,
211                    #0,
212                ),
213                "_rw_kafka_offset": (
214                    Varchar,
215                    #4,
216                ),
217                "_rw_kafka_partition": (
218                    Varchar,
219                    #3,
220                ),
221                "_rw_kafka_timestamp": (
222                    Timestamptz,
223                    #2,
224                ),
225                "v1": (
226                    Int32,
227                    #1,
228                ),
229                "v2": (
230                    Varchar,
231                    #5,
232                ),
233            }
234        "#]]
235        .assert_debug_eq(&altered_columns);
236
237        // Check version
238        assert_eq!(source.version + 1, altered_source.version);
239
240        // Check definition
241        expect_test::expect!["CREATE SOURCE s_shared (v1 INT, v2 CHARACTER VARYING) WITH (connector = 'kafka', topic = 'abc', properties.bootstrap.server = 'localhost:29092') FORMAT PLAIN ENCODE JSON"].assert_eq(&altered_source.definition);
242
243        let sql = "alter source s add column v2 varchar;";
244        frontend.run_sql(sql).await.unwrap();
245
246        let altered_source = get_source("s");
247        let altered_columns: BTreeMap<_, _> = altered_source
248            .columns
249            .iter()
250            .map(|col| (col.name(), (col.data_type().clone(), col.column_id())))
251            .collect();
252
253        // Check the new column is added.
254        // Check the old columns and IDs are not changed.
255        expect_test::expect![[r#"
256            {
257                "_row_id": (
258                    Serial,
259                    #0,
260                ),
261                "_rw_kafka_timestamp": (
262                    Timestamptz,
263                    #2,
264                ),
265                "v1": (
266                    Int32,
267                    #1,
268                ),
269                "v2": (
270                    Varchar,
271                    #3,
272                ),
273            }
274        "#]]
275        .assert_debug_eq(&altered_columns);
276
277        // Check version
278        assert_eq!(source.version + 1, altered_source.version);
279
280        // Check definition
281        expect_test::expect!["CREATE SOURCE s (v1 INT, v2 CHARACTER VARYING) WITH (connector = 'kafka', topic = 'abc', properties.bootstrap.server = 'localhost:29092') FORMAT PLAIN ENCODE JSON"].assert_eq(&altered_source.definition);
282    }
283}