1use 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
24pub struct FlightServer {
26 data_store: Arc<RwLock<HashMap<String, Arc<RecordBatch>>>>,
28 auth_tokens: Arc<RwLock<HashMap<String, String>>>,
30 enable_auth: bool,
32 task_executor: Option<Arc<Worker>>,
35}
36
37impl FlightServer {
38 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 pub fn with_auth(mut self) -> Self {
50 self.enable_auth = true;
51 self
52 }
53
54 pub fn with_worker(mut self, worker: Arc<Worker>) -> Self {
62 self.task_executor = Some(worker);
63 self
64 }
65
66 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 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 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 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 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 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 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 pub fn into_service(self) -> FlightServiceServer<Self> {
162 FlightServiceServer::new(self)
163 }
164
165 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 debug!("Handshake request received");
240
241 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 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 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 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 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 while let Some(data_result) = stream.next().await {
396 flight_data_vec.push(data_result?);
397 }
398
399 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 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 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 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 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 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 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: None,
613 progress: Some(1.0),
614 expiration_time: None,
615 }))
616 } else {
617 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 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 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 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 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 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 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 let worker = Arc::new(Worker::new(WorkerConfig::new("w-exec".to_string())));
839 let server = FlightServer::new().with_worker(worker);
840
841 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 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 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}