1use 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
22pub struct FlightServer {
24 data_store: Arc<RwLock<HashMap<String, Arc<RecordBatch>>>>,
26 auth_tokens: Arc<RwLock<HashMap<String, String>>>,
28 enable_auth: bool,
30}
31
32impl FlightServer {
33 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 pub fn with_auth(mut self) -> Self {
44 self.enable_auth = true;
45 self
46 }
47
48 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 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 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 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 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 pub fn into_service(self) -> FlightServiceServer<Self> {
102 FlightServiceServer::new(self)
103 }
104
105 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 debug!("Handshake request received");
180
181 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 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 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 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 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 while let Some(data_result) = stream.next().await {
336 flight_data_vec.push(data_result?);
337 }
338
339 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 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 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 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 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 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 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: None,
544 progress: Some(1.0),
545 expiration_time: None,
546 }))
547 } else {
548 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 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 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 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 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 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 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 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}