Skip to main content

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