Skip to main content

oxigdal_distributed/flight/
server.rs

1//! Arrow Flight server implementation for distributed data transfer.
2//!
3//! This module implements an Arrow Flight server that streams geospatial data
4//! between nodes using zero-copy transfers.
5
6use crate::error::{DistributedError, Result};
7use arrow::record_batch::RecordBatch;
8use arrow_flight::{
9    Action, ActionType, Criteria, Empty, FlightData, FlightDescriptor, FlightEndpoint, FlightInfo,
10    HandshakeRequest, HandshakeResponse, PutResult, SchemaAsIpc, SchemaResult, Ticket,
11    flight_service_server::{FlightService, FlightServiceServer},
12};
13use arrow_ipc::writer::IpcWriteOptions;
14use bytes::Bytes;
15use futures::{Stream, StreamExt, stream};
16use std::collections::HashMap;
17use std::pin::Pin;
18use std::sync::{Arc, RwLock};
19use tonic::{Request, Response, Streaming};
20use tracing::{debug, info};
21
22/// Flight server for serving geospatial data.
23pub struct FlightServer {
24    /// Stored data partitions (ticket -> RecordBatch).
25    data_store: Arc<RwLock<HashMap<String, Arc<RecordBatch>>>>,
26    /// Authentication tokens.
27    auth_tokens: Arc<RwLock<HashMap<String, String>>>,
28    /// Enable authentication.
29    enable_auth: bool,
30}
31
32impl FlightServer {
33    /// Create a new Flight server.
34    pub fn new() -> Self {
35        Self {
36            data_store: Arc::new(RwLock::new(HashMap::new())),
37            auth_tokens: Arc::new(RwLock::new(HashMap::new())),
38            enable_auth: false,
39        }
40    }
41
42    /// Enable authentication.
43    pub fn with_auth(mut self) -> Self {
44        self.enable_auth = true;
45        self
46    }
47
48    /// Store data with a ticket.
49    pub fn store_data(&self, ticket: String, data: Arc<RecordBatch>) -> Result<()> {
50        let mut store = self
51            .data_store
52            .write()
53            .map_err(|_| DistributedError::flight_rpc("Failed to acquire data store lock"))?;
54
55        store.insert(ticket, data);
56        Ok(())
57    }
58
59    /// Retrieve data by ticket.
60    pub fn get_data(&self, ticket: &str) -> Result<Option<Arc<RecordBatch>>> {
61        let store = self
62            .data_store
63            .read()
64            .map_err(|_| DistributedError::flight_rpc("Failed to acquire data store lock"))?;
65
66        Ok(store.get(ticket).cloned())
67    }
68
69    /// Remove data by ticket.
70    pub fn remove_data(&self, ticket: &str) -> Result<Option<Arc<RecordBatch>>> {
71        let mut store = self
72            .data_store
73            .write()
74            .map_err(|_| DistributedError::flight_rpc("Failed to acquire data store lock"))?;
75
76        Ok(store.remove(ticket))
77    }
78
79    /// List all available tickets.
80    pub fn list_tickets(&self) -> Result<Vec<String>> {
81        let store = self
82            .data_store
83            .read()
84            .map_err(|_| DistributedError::flight_rpc("Failed to acquire data store lock"))?;
85
86        Ok(store.keys().cloned().collect())
87    }
88
89    /// Add authentication token.
90    pub fn add_auth_token(&self, token: String, user: String) -> Result<()> {
91        let mut tokens = self
92            .auth_tokens
93            .write()
94            .map_err(|_| DistributedError::authentication("Failed to acquire auth tokens lock"))?;
95
96        tokens.insert(token, user);
97        Ok(())
98    }
99
100    /// Convert to tonic service.
101    pub fn into_service(self) -> FlightServiceServer<Self> {
102        FlightServiceServer::new(self)
103    }
104
105    /// Enforce bearer-token authentication for an incoming request.
106    ///
107    /// When authentication is disabled this is a no-op. Otherwise the request must
108    /// carry an `authorization: Bearer <token>` metadata header whose token is a
109    /// registered key in [`Self::add_auth_token`]. Any missing, malformed, or unknown
110    /// token is rejected with [`tonic::Status::unauthenticated`]; a poisoned token
111    /// lock is reported as [`tonic::Status::internal`].
112    fn check_auth<T>(&self, request: &Request<T>) -> std::result::Result<(), tonic::Status> {
113        if !self.enable_auth {
114            return Ok(());
115        }
116
117        let header = request
118            .metadata()
119            .get("authorization")
120            .ok_or_else(|| tonic::Status::unauthenticated("Missing authorization header"))?;
121
122        let value = header
123            .to_str()
124            .map_err(|_| tonic::Status::unauthenticated("Invalid authorization header encoding"))?;
125
126        let token = value
127            .strip_prefix("Bearer ")
128            .or_else(|| value.strip_prefix("bearer "))
129            .map(str::trim)
130            .filter(|t| !t.is_empty())
131            .ok_or_else(|| {
132                tonic::Status::unauthenticated("Authorization header must be a Bearer token")
133            })?;
134
135        let tokens = self
136            .auth_tokens
137            .read()
138            .map_err(|_| tonic::Status::internal("Failed to acquire auth tokens lock"))?;
139
140        if tokens.contains_key(token) {
141            Ok(())
142        } else {
143            Err(tonic::Status::unauthenticated("Invalid or unknown token"))
144        }
145    }
146}
147
148impl Default for FlightServer {
149    fn default() -> Self {
150        Self::new()
151    }
152}
153
154#[tonic::async_trait]
155impl FlightService for FlightServer {
156    type HandshakeStream =
157        Pin<Box<dyn Stream<Item = std::result::Result<HandshakeResponse, tonic::Status>> + Send>>;
158    type ListFlightsStream =
159        Pin<Box<dyn Stream<Item = std::result::Result<FlightInfo, tonic::Status>> + Send>>;
160    type DoGetStream =
161        Pin<Box<dyn Stream<Item = std::result::Result<FlightData, tonic::Status>> + Send>>;
162    type DoPutStream =
163        Pin<Box<dyn Stream<Item = std::result::Result<PutResult, tonic::Status>> + Send>>;
164    type DoActionStream = Pin<
165        Box<dyn Stream<Item = std::result::Result<arrow_flight::Result, tonic::Status>> + Send>,
166    >;
167    type ListActionsStream =
168        Pin<Box<dyn Stream<Item = std::result::Result<ActionType, tonic::Status>> + Send>>;
169    type DoExchangeStream =
170        Pin<Box<dyn Stream<Item = std::result::Result<FlightData, tonic::Status>> + Send>>;
171
172    async fn handshake(
173        &self,
174        _request: Request<Streaming<HandshakeRequest>>,
175    ) -> std::result::Result<Response<Self::HandshakeStream>, tonic::Status> {
176        // Handshake is deliberately left un-gated: it is the credential-exchange entry
177        // point of the Flight protocol and returns no data-store contents. Every other
178        // RPC method enforces `check_auth` before touching the data store.
179        debug!("Handshake request received");
180
181        // Simple handshake - just acknowledge
182        let response = HandshakeResponse {
183            protocol_version: 0,
184            payload: Bytes::new(),
185        };
186
187        let stream = stream::once(async { Ok(response) });
188        Ok(Response::new(Box::pin(stream)))
189    }
190
191    async fn list_flights(
192        &self,
193        request: Request<Criteria>,
194    ) -> std::result::Result<Response<Self::ListFlightsStream>, tonic::Status> {
195        self.check_auth(&request)?;
196        debug!("List flights request received");
197
198        // Return empty stream - we don't support flight listing yet
199        let stream = stream::empty();
200        Ok(Response::new(Box::pin(stream)))
201    }
202
203    async fn get_flight_info(
204        &self,
205        request: Request<FlightDescriptor>,
206    ) -> std::result::Result<Response<FlightInfo>, tonic::Status> {
207        self.check_auth(&request)?;
208        let descriptor = request.into_inner();
209        debug!("Get flight info request: {:?}", descriptor);
210
211        // Resolve ticket key from descriptor: prefer first path segment, fall back to cmd bytes.
212        let ticket_key = if !descriptor.path.is_empty() {
213            descriptor.path[0].clone()
214        } else if !descriptor.cmd.is_empty() {
215            String::from_utf8(descriptor.cmd.to_vec())
216                .map_err(|e| tonic::Status::invalid_argument(format!("Invalid cmd: {}", e)))?
217        } else {
218            return Err(tonic::Status::invalid_argument(
219                "FlightDescriptor must have a path or cmd",
220            ));
221        };
222
223        let data = self
224            .get_data(&ticket_key)
225            .map_err(|e| tonic::Status::internal(e.to_string()))?
226            .ok_or_else(|| tonic::Status::not_found(format!("Flight not found: {}", ticket_key)))?;
227
228        let schema = data.schema();
229        let ipc_opts = IpcWriteOptions::default();
230        let schema_bytes: arrow_flight::IpcMessage = SchemaAsIpc::new(schema.as_ref(), &ipc_opts)
231            .try_into()
232            .map_err(|e: arrow_schema::ArrowError| {
233                tonic::Status::internal(format!("Schema encode error: {}", e))
234            })?;
235
236        let endpoint = FlightEndpoint {
237            ticket: Some(Ticket {
238                ticket: Bytes::from(ticket_key),
239            }),
240            location: vec![],
241            expiration_time: None,
242            app_metadata: Bytes::new(),
243        };
244
245        let flight_info = FlightInfo {
246            schema: schema_bytes.0,
247            flight_descriptor: Some(descriptor),
248            endpoint: vec![endpoint],
249            total_records: data.num_rows() as i64,
250            total_bytes: -1,
251            ordered: false,
252            app_metadata: Bytes::new(),
253        };
254
255        Ok(Response::new(flight_info))
256    }
257
258    async fn get_schema(
259        &self,
260        request: Request<FlightDescriptor>,
261    ) -> std::result::Result<Response<SchemaResult>, tonic::Status> {
262        self.check_auth(&request)?;
263        let descriptor = request.into_inner();
264        debug!("Get schema request received");
265
266        let ticket_key = if !descriptor.path.is_empty() {
267            descriptor.path[0].clone()
268        } else if !descriptor.cmd.is_empty() {
269            String::from_utf8(descriptor.cmd.to_vec())
270                .map_err(|e| tonic::Status::invalid_argument(format!("Invalid cmd: {}", e)))?
271        } else {
272            return Err(tonic::Status::invalid_argument(
273                "FlightDescriptor must have a path or cmd",
274            ));
275        };
276
277        let data = self
278            .get_data(&ticket_key)
279            .map_err(|e| tonic::Status::internal(e.to_string()))?
280            .ok_or_else(|| tonic::Status::not_found(format!("Flight not found: {}", ticket_key)))?;
281
282        let schema = data.schema();
283        let ipc_opts = IpcWriteOptions::default();
284        let schema_result: SchemaResult = SchemaAsIpc::new(schema.as_ref(), &ipc_opts)
285            .try_into()
286            .map_err(|e: arrow_schema::ArrowError| {
287            tonic::Status::internal(format!("Schema encode error: {}", e))
288        })?;
289
290        Ok(Response::new(schema_result))
291    }
292
293    async fn do_get(
294        &self,
295        request: Request<Ticket>,
296    ) -> std::result::Result<Response<Self::DoGetStream>, tonic::Status> {
297        self.check_auth(&request)?;
298        let ticket = request.into_inner();
299        let ticket_str = String::from_utf8(ticket.ticket.to_vec())
300            .map_err(|e| tonic::Status::invalid_argument(format!("Invalid ticket: {}", e)))?;
301
302        info!("DoGet request for ticket: {}", ticket_str);
303
304        // Retrieve data
305        let data = self
306            .get_data(&ticket_str)
307            .map_err(|e| tonic::Status::internal(e.to_string()))?
308            .ok_or_else(|| tonic::Status::not_found(format!("Ticket not found: {}", ticket_str)))?;
309
310        // Convert RecordBatch to FlightData stream
311        let flight_data_vec = arrow_flight::utils::batches_to_flight_data(
312            data.schema().as_ref(),
313            vec![(*data).clone()],
314        )
315        .map_err(|e| tonic::Status::internal(format!("Failed to encode batches: {}", e)))?
316        .into_iter()
317        .map(Ok)
318        .collect::<Vec<_>>();
319
320        let stream = stream::iter(flight_data_vec);
321        Ok(Response::new(Box::pin(stream)))
322    }
323
324    async fn do_put(
325        &self,
326        request: Request<Streaming<FlightData>>,
327    ) -> std::result::Result<Response<Self::DoPutStream>, tonic::Status> {
328        self.check_auth(&request)?;
329        debug!("DoPut request received");
330
331        let mut stream = request.into_inner();
332        let mut flight_data_vec = Vec::new();
333
334        // Collect all FlightData messages
335        while let Some(data_result) = stream.next().await {
336            flight_data_vec.push(data_result?);
337        }
338
339        // Convert FlightData to RecordBatches
340        let batches = arrow_flight::utils::flight_data_to_batches(&flight_data_vec)
341            .map_err(|e| tonic::Status::internal(format!("Failed to decode batches: {}", e)))?;
342
343        info!("DoPut received {} batches", batches.len());
344
345        // Store batches (using a generated ticket)
346        for (i, batch) in batches.into_iter().enumerate() {
347            let ticket = format!("uploaded_{}", i);
348            self.store_data(ticket, Arc::new(batch))
349                .map_err(|e| tonic::Status::internal(e.to_string()))?;
350        }
351
352        // Return success
353        let result = PutResult {
354            app_metadata: Bytes::new(),
355        };
356
357        let stream = stream::once(async { Ok(result) });
358        Ok(Response::new(Box::pin(stream)))
359    }
360
361    async fn do_action(
362        &self,
363        request: Request<Action>,
364    ) -> std::result::Result<Response<Self::DoActionStream>, tonic::Status> {
365        self.check_auth(&request)?;
366        let action = request.into_inner();
367        info!("DoAction request: {}", action.r#type);
368
369        match action.r#type.as_str() {
370            "list_tickets" => {
371                let tickets = self
372                    .list_tickets()
373                    .map_err(|e| tonic::Status::internal(e.to_string()))?;
374
375                let result = arrow_flight::Result {
376                    body: serde_json::to_vec(&tickets)
377                        .map_err(|e| {
378                            tonic::Status::internal(format!("Serialization error: {}", e))
379                        })?
380                        .into(),
381                };
382
383                let stream = stream::once(async { Ok(result) });
384                Ok(Response::new(Box::pin(stream)))
385            }
386            "remove_ticket" => {
387                let ticket = String::from_utf8(action.body.to_vec()).map_err(|e| {
388                    tonic::Status::invalid_argument(format!("Invalid ticket: {}", e))
389                })?;
390
391                self.remove_data(&ticket)
392                    .map_err(|e| tonic::Status::internal(e.to_string()))?;
393
394                let result = arrow_flight::Result {
395                    body: Bytes::from("removed"),
396                };
397
398                let stream = stream::once(async { Ok(result) });
399                Ok(Response::new(Box::pin(stream)))
400            }
401            "list_actions" => {
402                // Reflection action: return JSON array of supported action descriptors.
403                let actions = vec![
404                    serde_json::json!({"type": "list_tickets", "description": "List all available tickets"}),
405                    serde_json::json!({"type": "remove_ticket", "description": "Remove a ticket from the server"}),
406                    serde_json::json!({"type": "list_actions", "description": "List all supported actions (reflection)"}),
407                    serde_json::json!({"type": "ping", "description": "Health check — returns 'pong'"}),
408                ];
409
410                let result = arrow_flight::Result {
411                    body: serde_json::to_vec(&actions)
412                        .map_err(|e| {
413                            tonic::Status::internal(format!("Serialization error: {}", e))
414                        })?
415                        .into(),
416                };
417
418                let stream = stream::once(async { Ok(result) });
419                Ok(Response::new(Box::pin(stream)))
420            }
421            "ping" => {
422                let result = arrow_flight::Result {
423                    body: Bytes::from_static(b"pong"),
424                };
425                let stream = stream::once(async { Ok(result) });
426                Ok(Response::new(Box::pin(stream)))
427            }
428            _ => Err(tonic::Status::unimplemented(format!(
429                "Action not implemented: {}",
430                action.r#type
431            ))),
432        }
433    }
434
435    async fn list_actions(
436        &self,
437        request: Request<Empty>,
438    ) -> std::result::Result<Response<Self::ListActionsStream>, tonic::Status> {
439        self.check_auth(&request)?;
440        debug!("List actions request received");
441
442        let actions = vec![
443            ActionType {
444                r#type: "list_tickets".to_string(),
445                description: "List all available tickets".to_string(),
446            },
447            ActionType {
448                r#type: "remove_ticket".to_string(),
449                description: "Remove a ticket from the server".to_string(),
450            },
451        ];
452
453        let stream = stream::iter(actions.into_iter().map(Ok));
454        Ok(Response::new(Box::pin(stream)))
455    }
456
457    async fn do_exchange(
458        &self,
459        request: Request<Streaming<FlightData>>,
460    ) -> std::result::Result<Response<Self::DoExchangeStream>, tonic::Status> {
461        self.check_auth(&request)?;
462        debug!("DoExchange request received — echo/passthrough mode");
463
464        let mut incoming = request.into_inner();
465        let mut echo_items: Vec<std::result::Result<FlightData, tonic::Status>> = Vec::new();
466
467        while let Some(item) = incoming.next().await {
468            match item {
469                Ok(flight_data) => {
470                    echo_items.push(Ok(flight_data));
471                }
472                Err(status) => {
473                    // Propagate the first error as the terminal item.
474                    echo_items.push(Err(status));
475                    break;
476                }
477            }
478        }
479
480        info!("DoExchange echoing {} items", echo_items.len());
481        let stream = stream::iter(echo_items);
482        Ok(Response::new(Box::pin(stream)))
483    }
484
485    async fn poll_flight_info(
486        &self,
487        request: Request<FlightDescriptor>,
488    ) -> std::result::Result<Response<arrow_flight::PollInfo>, tonic::Status> {
489        self.check_auth(&request)?;
490        let descriptor = request.into_inner();
491        debug!("Poll flight info request received");
492
493        // Resolve the ticket key from the descriptor (same logic as get_flight_info).
494        let ticket_key = if !descriptor.path.is_empty() {
495            descriptor.path[0].clone()
496        } else if !descriptor.cmd.is_empty() {
497            String::from_utf8(descriptor.cmd.to_vec())
498                .map_err(|e| tonic::Status::invalid_argument(format!("Invalid cmd: {}", e)))?
499        } else {
500            return Err(tonic::Status::invalid_argument(
501                "FlightDescriptor must have a path or cmd",
502            ));
503        };
504
505        let data_opt = self
506            .get_data(&ticket_key)
507            .map_err(|e| tonic::Status::internal(e.to_string()))?;
508
509        // If data is ready, return complete FlightInfo (no pending descriptor, progress = 1.0).
510        // If data is still pending, echo the descriptor back (client should retry), progress = None.
511        if let Some(data) = data_opt {
512            let schema = data.schema();
513            let ipc_opts = IpcWriteOptions::default();
514            let schema_bytes: arrow_flight::IpcMessage =
515                SchemaAsIpc::new(schema.as_ref(), &ipc_opts)
516                    .try_into()
517                    .map_err(|e: arrow_schema::ArrowError| {
518                        tonic::Status::internal(format!("Schema encode error: {}", e))
519                    })?;
520
521            let endpoint = FlightEndpoint {
522                ticket: Some(Ticket {
523                    ticket: Bytes::from(ticket_key),
524                }),
525                location: vec![],
526                expiration_time: None,
527                app_metadata: Bytes::new(),
528            };
529
530            let flight_info = FlightInfo {
531                schema: schema_bytes.0,
532                flight_descriptor: Some(descriptor),
533                endpoint: vec![endpoint],
534                total_records: data.num_rows() as i64,
535                total_bytes: -1,
536                ordered: false,
537                app_metadata: Bytes::new(),
538            };
539
540            Ok(Response::new(arrow_flight::PollInfo {
541                info: Some(flight_info),
542                // flight_descriptor is None -> indicates the query is complete.
543                flight_descriptor: None,
544                progress: Some(1.0),
545                expiration_time: None,
546            }))
547        } else {
548            // Data not yet ready; client should poll again using the same descriptor.
549            Ok(Response::new(arrow_flight::PollInfo {
550                info: None,
551                flight_descriptor: Some(descriptor),
552                progress: None,
553                expiration_time: None,
554            }))
555        }
556    }
557}
558
559#[cfg(test)]
560mod tests {
561    use super::*;
562    use arrow::array::Int32Array;
563    use arrow::datatypes::{DataType, Field, Schema};
564
565    fn create_test_batch() -> std::result::Result<Arc<RecordBatch>, Box<dyn std::error::Error>> {
566        let schema = Arc::new(Schema::new(vec![Field::new(
567            "value",
568            DataType::Int32,
569            false,
570        )]));
571
572        let array = Int32Array::from(vec![1, 2, 3, 4, 5]);
573
574        Ok(Arc::new(RecordBatch::try_new(
575            schema,
576            vec![Arc::new(array)],
577        )?))
578    }
579
580    #[test]
581    fn test_server_creation() {
582        let server = FlightServer::new();
583        assert!(!server.enable_auth);
584    }
585
586    #[test]
587    fn test_store_and_retrieve_data() -> std::result::Result<(), Box<dyn std::error::Error>> {
588        let server = FlightServer::new();
589        let batch = create_test_batch()?;
590
591        server.store_data("test_ticket".to_string(), batch.clone())?;
592
593        let retrieved = server
594            .get_data("test_ticket")?
595            .ok_or_else(|| Box::<dyn std::error::Error>::from("should exist"))?;
596
597        assert_eq!(retrieved.num_rows(), batch.num_rows());
598        Ok(())
599    }
600
601    #[test]
602    fn test_remove_data() -> std::result::Result<(), Box<dyn std::error::Error>> {
603        let server = FlightServer::new();
604        let batch = create_test_batch()?;
605
606        server.store_data("test_ticket".to_string(), batch)?;
607
608        let removed = server
609            .remove_data("test_ticket")?
610            .ok_or_else(|| Box::<dyn std::error::Error>::from("should exist"))?;
611
612        assert_eq!(removed.num_rows(), 5);
613
614        let retrieved = server.get_data("test_ticket")?;
615        assert!(retrieved.is_none());
616        Ok(())
617    }
618
619    #[test]
620    fn test_list_tickets() -> std::result::Result<(), Box<dyn std::error::Error>> {
621        let server = FlightServer::new();
622
623        server.store_data("ticket1".to_string(), create_test_batch()?)?;
624        server.store_data("ticket2".to_string(), create_test_batch()?)?;
625
626        let tickets = server.list_tickets()?;
627        assert_eq!(tickets.len(), 2);
628        assert!(tickets.contains(&"ticket1".to_string()));
629        assert!(tickets.contains(&"ticket2".to_string()));
630        Ok(())
631    }
632
633    #[test]
634    fn test_authentication() -> std::result::Result<(), Box<dyn std::error::Error>> {
635        let server = FlightServer::new().with_auth();
636        assert!(server.enable_auth);
637
638        server.add_auth_token("token123".to_string(), "user1".to_string())?;
639
640        // Verify token exists via auth_tokens (verify_token method not exposed)
641        assert!(
642            server
643                .auth_tokens
644                .read()
645                .map_err(|e| Box::<dyn std::error::Error>::from(format!("lock poisoned: {}", e)))?
646                .contains_key("token123")
647        );
648        assert!(
649            !server
650                .auth_tokens
651                .read()
652                .map_err(|e| Box::<dyn std::error::Error>::from(format!("lock poisoned: {}", e)))?
653                .contains_key("invalid")
654        );
655        Ok(())
656    }
657
658    fn bearer_request<T>(inner: T, token: &str) -> Request<T> {
659        let mut request = Request::new(inner);
660        let value = format!("Bearer {}", token);
661        if let Ok(meta_value) = value.parse::<tonic::metadata::MetadataValue<_>>() {
662            request.metadata_mut().insert("authorization", meta_value);
663        }
664        request
665    }
666
667    #[tokio::test]
668    async fn test_do_get_rejects_unauthenticated()
669    -> std::result::Result<(), Box<dyn std::error::Error>> {
670        let server = FlightServer::new().with_auth();
671        server.add_auth_token("token123".to_string(), "user1".to_string())?;
672        server.store_data("t1".to_string(), create_test_batch()?)?;
673
674        // No authorization header -> unauthenticated.
675        let req = Request::new(Ticket {
676            ticket: Bytes::from("t1"),
677        });
678        let result = FlightService::do_get(&server, req).await;
679        assert!(result.is_err());
680        assert_eq!(
681            result.err().map(|s| s.code()),
682            Some(tonic::Code::Unauthenticated)
683        );
684
685        // Wrong token -> unauthenticated.
686        let bad = bearer_request(
687            Ticket {
688                ticket: Bytes::from("t1"),
689            },
690            "wrong",
691        );
692        let result = FlightService::do_get(&server, bad).await;
693        assert_eq!(
694            result.err().map(|s| s.code()),
695            Some(tonic::Code::Unauthenticated)
696        );
697
698        // Valid token -> succeeds.
699        let good = bearer_request(
700            Ticket {
701                ticket: Bytes::from("t1"),
702            },
703            "token123",
704        );
705        let result = FlightService::do_get(&server, good).await;
706        assert!(result.is_ok());
707        Ok(())
708    }
709
710    #[tokio::test]
711    async fn test_do_action_remove_requires_auth()
712    -> std::result::Result<(), Box<dyn std::error::Error>> {
713        let server = FlightServer::new().with_auth();
714        server.add_auth_token("token123".to_string(), "user1".to_string())?;
715        server.store_data("t1".to_string(), create_test_batch()?)?;
716
717        // Unauthenticated remove_ticket must be rejected and must NOT delete data.
718        let action = Action {
719            r#type: "remove_ticket".to_string(),
720            body: Bytes::from("t1"),
721        };
722        let result = FlightService::do_action(&server, Request::new(action)).await;
723        assert_eq!(
724            result.err().map(|s| s.code()),
725            Some(tonic::Code::Unauthenticated)
726        );
727        assert!(server.get_data("t1")?.is_some(), "data must not be removed");
728
729        // Authenticated remove succeeds.
730        let action = Action {
731            r#type: "remove_ticket".to_string(),
732            body: Bytes::from("t1"),
733        };
734        let result = FlightService::do_action(&server, bearer_request(action, "token123")).await;
735        assert!(result.is_ok());
736        assert!(server.get_data("t1")?.is_none());
737        Ok(())
738    }
739
740    #[tokio::test]
741    async fn test_auth_disabled_allows_access()
742    -> std::result::Result<(), Box<dyn std::error::Error>> {
743        // Without with_auth(), requests without any token must still succeed.
744        let server = FlightServer::new();
745        server.store_data("t1".to_string(), create_test_batch()?)?;
746
747        let req = Request::new(Ticket {
748            ticket: Bytes::from("t1"),
749        });
750        let result = FlightService::do_get(&server, req).await;
751        assert!(result.is_ok());
752        Ok(())
753    }
754}