1use std::sync::Arc;
2use std::sync::atomic::{AtomicUsize, Ordering};
3use std::time::{Duration, Instant};
4
5use async_trait::async_trait;
6use camel_bridge::download::{default_cache_dir_for_spec, ensure_binary_for_spec};
7use camel_bridge::health::wait_for_health;
8use camel_bridge::process::{BridgeError, BridgeProcess, BridgeProcessConfig};
9use camel_bridge::reconnect::BridgeReconnectHandler;
10use camel_bridge::spec::XML_BRIDGE;
11use camel_component_api::RuntimeObservability;
12use dashmap::DashMap;
13use sha2::{Digest, Sha256};
14use std::sync::OnceLock;
15use tokio::sync::{Mutex, RwLock, watch};
16use tonic::Code;
17use tonic::transport::Channel;
18use tracing::error;
19
20use crate::error::ValidatorError;
21use crate::proto;
22use crate::proto::{
23 HealthCheckRequest, RegisterSchemaRequest, RegisterSchemaResponse, ValidateResponse,
24 ValidateWithRequest,
25};
26
27pub type SchemaId = String;
28
29pub(crate) const XSD_BRIDGE_MAX_DECODING_MESSAGE_SIZE: usize = 17 * 1024 * 1024;
41
42pub(crate) fn xsd_bridge_decode_limit() -> usize {
43 XSD_BRIDGE_MAX_DECODING_MESSAGE_SIZE
44}
45
46pub(crate) fn xsd_bridge_client(
53 channel: Channel,
54) -> proto::xsd_validator_client::XsdValidatorClient<Channel> {
55 proto::xsd_validator_client::XsdValidatorClient::new(channel)
56 .max_decoding_message_size(xsd_bridge_decode_limit())
57}
58
59#[derive(Debug, Clone)]
60pub enum BridgeState {
61 Starting,
62 Ready { channel: Channel },
63 Degraded(String),
64 Restarting { attempt: u32, next_at: Instant },
65 Stopped,
66}
67
68pub struct XmlBridgeSlot {
69 pub state_rx: watch::Receiver<BridgeState>,
70 pub(crate) state_tx: watch::Sender<BridgeState>,
71 pub process: Arc<Mutex<Option<BridgeProcess>>>,
72}
73
74impl std::fmt::Debug for XmlBridgeSlot {
75 fn fmt(&self, f: &mut std::fmt::Formatter<'_>) -> std::fmt::Result {
76 f.debug_struct("XmlBridgeSlot").finish()
77 }
78}
79
80#[async_trait]
81pub trait XsdBridge: Send + Sync {
82 async fn register(&self, xsd_bytes: Vec<u8>) -> Result<SchemaId, ValidatorError>;
83 async fn validate(&self, schema_id: &str, doc_bytes: Vec<u8>) -> Result<(), ValidatorError>;
84}
85
86#[async_trait]
87trait XsdBridgeRpc: Send + Sync {
88 async fn register_schema(
89 &self,
90 channel: Channel,
91 request: RegisterSchemaRequest,
92 ) -> Result<RegisterSchemaResponse, ValidatorError>;
93
94 async fn validate_with(
95 &self,
96 channel: Channel,
97 request: ValidateWithRequest,
98 ) -> Result<ValidateResponse, ValidatorError>;
99}
100
101#[derive(Debug)]
102struct GrpcXsdBridgeRpc;
103
104#[async_trait]
105impl XsdBridgeRpc for GrpcXsdBridgeRpc {
106 async fn register_schema(
107 &self,
108 channel: Channel,
109 request: RegisterSchemaRequest,
110 ) -> Result<RegisterSchemaResponse, ValidatorError> {
111 let mut client = xsd_bridge_client(channel);
112 let response = client.register_schema(request).await.map_err(|e| {
113 ValidatorError::transport_with_source("xml-bridge register_schema RPC failed", e)
114 })?;
115 Ok(response.into_inner())
116 }
117
118 async fn validate_with(
119 &self,
120 channel: Channel,
121 request: ValidateWithRequest,
122 ) -> Result<ValidateResponse, ValidatorError> {
123 let mut client = xsd_bridge_client(channel);
124 let response = client.validate_with(request).await.map_err(|e| {
125 ValidatorError::transport_with_source("xml-bridge validate_with RPC failed", e)
126 })?;
127 Ok(response.into_inner())
128 }
129}
130
131#[derive(Clone)]
132pub struct XsdBridgeBackend {
133 channel: Arc<RwLock<Option<Channel>>>,
134 schemas: Arc<DashMap<SchemaId, Vec<u8>>>,
135 schema_cache_max_entries: Arc<AtomicUsize>,
136 slot: Arc<XmlBridgeSlot>,
137 rpc: Arc<dyn XsdBridgeRpc>,
138 start_lock: Arc<Mutex<()>>,
139 bridge_version: String,
140 bridge_cache_dir: std::path::PathBuf,
141 bridge_start_timeout_ms: u64,
142 observability: Arc<OnceLock<(Arc<dyn RuntimeObservability>, String)>>,
143}
144
145impl std::fmt::Debug for XsdBridgeBackend {
146 fn fmt(&self, f: &mut std::fmt::Formatter<'_>) -> std::fmt::Result {
147 f.debug_struct("XsdBridgeBackend")
148 .field("bridge_version", &self.bridge_version)
149 .field("bridge_cache_dir", &self.bridge_cache_dir)
150 .finish()
151 }
152}
153
154impl XsdBridgeBackend {
155 pub fn new() -> Self {
156 let (state_tx, state_rx) = watch::channel(BridgeState::Stopped);
157 let slot = Arc::new(XmlBridgeSlot {
158 state_rx,
159 state_tx,
160 process: Arc::new(Mutex::new(None)),
161 });
162
163 Self {
164 channel: Arc::new(RwLock::new(None)),
165 schemas: Arc::new(DashMap::new()),
166 schema_cache_max_entries: Arc::new(AtomicUsize::new(
167 crate::config::DEFAULT_SCHEMA_CACHE_MAX_ENTRIES,
168 )),
169 slot,
170 rpc: Arc::new(GrpcXsdBridgeRpc),
171 start_lock: Arc::new(Mutex::new(())),
172 bridge_version: crate::BRIDGE_VERSION.to_string(),
173 bridge_cache_dir: default_cache_dir_for_spec(&XML_BRIDGE),
174 bridge_start_timeout_ms: 30_000,
175 observability: Arc::new(OnceLock::new()),
176 }
177 }
178
179 #[cfg(test)]
180 fn for_test(rpc: Arc<dyn XsdBridgeRpc>, channel: Channel) -> Self {
181 let (state_tx, state_rx) = watch::channel(BridgeState::Ready {
182 channel: channel.clone(),
183 });
184 let slot = Arc::new(XmlBridgeSlot {
185 state_rx,
186 state_tx,
187 process: Arc::new(Mutex::new(None)),
188 });
189 Self {
190 channel: Arc::new(RwLock::new(Some(channel))),
191 schemas: Arc::new(DashMap::new()),
192 schema_cache_max_entries: Arc::new(AtomicUsize::new(
193 crate::config::DEFAULT_SCHEMA_CACHE_MAX_ENTRIES,
194 )),
195 slot,
196 rpc,
197 start_lock: Arc::new(Mutex::new(())),
198 bridge_version: crate::BRIDGE_VERSION.to_string(),
199 bridge_cache_dir: default_cache_dir_for_spec(&XML_BRIDGE),
200 bridge_start_timeout_ms: 30_000,
201 observability: Arc::new(OnceLock::new()),
202 }
203 }
204
205 pub fn set_observability(&self, runtime: Arc<dyn RuntimeObservability>, route_id: String) {
206 self.observability.set((runtime, route_id)).ok();
207 }
208
209 pub fn schema_id_for(xsd_bytes: &[u8]) -> SchemaId {
210 let mut hasher = Sha256::new();
211 hasher.update(xsd_bytes);
212 format!("xsd-{}", hex::encode(hasher.finalize()))
213 }
214
215 pub fn set_schema_cache_max_entries(&self, max_entries: usize) {
218 if self.schemas.len() > max_entries {
219 self.schemas.clear();
220 }
221 self.schema_cache_max_entries
222 .store(max_entries, Ordering::Relaxed);
223 }
224
225 async fn ensure_bridge_ready(&self) -> Result<Channel, ValidatorError> {
226 if let Some(ch) = self.channel.read().await.clone() {
227 return Ok(ch);
228 }
229
230 let _guard = self.start_lock.lock().await;
231 if let Some(ch) = self.channel.read().await.clone() {
232 return Ok(ch);
233 }
234
235 let _ = self.slot.state_tx.send(BridgeState::Starting);
236 let (process, channel) = self.start_bridge_process().await?;
237 {
238 let mut process_guard = self.slot.process.lock().await;
239 *process_guard = Some(process);
240 }
241 {
242 let mut ch_guard = self.channel.write().await;
243 *ch_guard = Some(channel.clone());
244 }
245 let _ = self.slot.state_tx.send(BridgeState::Ready {
246 channel: channel.clone(),
247 });
248 self.on_reconnect(&channel).map_err(|e| {
249 ValidatorError::transport_with_source("xml-bridge reconnect handler failed", e)
250 })?;
251
252 Ok(channel)
253 }
254
255 async fn restart_bridge(&self) -> Result<Channel, ValidatorError> {
256 let _guard = self.start_lock.lock().await;
257
258 let _ = self.slot.state_tx.send(BridgeState::Restarting {
259 attempt: 0,
260 next_at: Instant::now(),
261 });
262
263 let old_process = {
264 let mut process_guard = self.slot.process.lock().await;
265 process_guard.take()
266 };
267 if let Some(p) = old_process {
268 let _ = p.stop().await;
269 }
270
271 let (process, channel) = self.start_bridge_process().await?;
272 {
273 let mut process_guard = self.slot.process.lock().await;
274 *process_guard = Some(process);
275 }
276 {
277 let mut ch_guard = self.channel.write().await;
278 *ch_guard = Some(channel.clone());
279 }
280
281 self.on_reconnect(&channel).map_err(|e| {
282 ValidatorError::transport_with_source("xml-bridge reconnect handler failed", e)
283 })?;
284 let _ = self.slot.state_tx.send(BridgeState::Ready {
285 channel: channel.clone(),
286 });
287
288 Ok(channel)
289 }
290
291 async fn start_bridge_process(&self) -> Result<(BridgeProcess, Channel), ValidatorError> {
292 let binary_path =
293 ensure_binary_for_spec(&XML_BRIDGE, &self.bridge_version, &self.bridge_cache_dir)
294 .await
295 .map_err(|e| {
296 ValidatorError::endpoint(format!("XML bridge binary unavailable: {e}"))
297 })?;
298
299 let process_config = BridgeProcessConfig::xml(binary_path, self.bridge_start_timeout_ms);
300 let (process, channel) = BridgeProcess::start_and_connect(&process_config)
301 .await
302 .map_err(|e| ValidatorError::endpoint(format!("XML bridge start failed: {e}")))?;
303
304 wait_for_health(&channel, Duration::from_secs(10), |ch| {
305 let mut client = proto::health_client::HealthClient::new(ch);
306 async move {
307 let resp = client.check(HealthCheckRequest {}).await?;
308 Ok(resp.into_inner().status == "SERVING")
309 }
310 })
311 .await
312 .map_err(|e| ValidatorError::endpoint(format!("XML bridge health check failed: {e}")))?;
313
314 Ok((process, channel))
315 }
316
317 fn is_transport_error(msg: &str) -> bool {
318 msg.contains(&Code::Unavailable.to_string())
319 || msg.contains(&Code::Unknown.to_string())
320 || msg.contains("transport")
321 }
322
323 pub async fn shutdown(&self) {
324 let mut guard = self.slot.process.lock().await;
325 if let Some(p) = guard.take()
326 && let Err(e) = p.stop().await
327 {
328 tracing::warn!("Failed to stop XSD bridge process: {}", e);
329 }
330 }
331
332 async fn register_with_channel(
333 &self,
334 channel: Channel,
335 schema_id: SchemaId,
336 xsd_bytes: Vec<u8>,
337 ) -> Result<SchemaId, ValidatorError> {
338 let response = self
339 .rpc
340 .register_schema(
341 channel,
342 RegisterSchemaRequest {
343 schema_id: schema_id.clone(),
344 schema: xsd_bytes.clone(),
345 },
346 )
347 .await?;
348
349 if let Some(err) = response.error {
350 return Err(ValidatorError::from_bridge_error(&err));
351 }
352
353 let max_entries = self.schema_cache_max_entries.load(Ordering::Relaxed);
357 if self.schemas.len() >= max_entries && !self.schemas.contains_key(&schema_id) {
358 self.schemas.clear();
359 }
360
361 self.schemas.insert(schema_id.clone(), xsd_bytes);
362 Ok(schema_id)
363 }
364}
365
366impl BridgeReconnectHandler for XsdBridgeBackend {
367 fn on_reconnect(&self, _channel: &Channel) -> Result<(), BridgeError> {
368 let this = self.clone();
369 tokio::spawn(async move {
370 let Some(channel) = this.channel.read().await.clone() else {
371 return;
372 };
373
374 let schemas: Vec<(SchemaId, Vec<u8>)> = this
375 .schemas
376 .iter()
377 .map(|entry| (entry.key().clone(), entry.value().clone()))
378 .collect();
379
380 let observability = this
381 .observability
382 .get()
383 .map(|(rt, rid)| (Arc::clone(rt), rid.clone()));
384
385 for (schema_id, schema_bytes) in schemas {
386 if let Err(e) = this
387 .rpc
388 .register_schema(
389 channel.clone(),
390 RegisterSchemaRequest {
391 schema_id: schema_id.clone(),
392 schema: schema_bytes,
393 },
394 )
395 .await
396 {
397 if let Some((ref rt, ref rid)) = observability {
398 rt.metrics()
399 .increment_errors(rid, "e:validator:reconnect-reseed");
400 }
401 error!(schema_id = %schema_id, error = %e, "re-seed schema failed after reconnect");
403 }
404 }
405 });
406 Ok(())
407 }
408}
409
410#[async_trait]
411impl XsdBridge for XsdBridgeBackend {
412 async fn register(&self, xsd_bytes: Vec<u8>) -> Result<SchemaId, ValidatorError> {
413 let schema_id = Self::schema_id_for(&xsd_bytes);
414 if self.schemas.contains_key(&schema_id) {
415 return Ok(schema_id);
416 }
417
418 let channel = self.ensure_bridge_ready().await?;
419 match self
420 .register_with_channel(channel.clone(), schema_id.clone(), xsd_bytes.clone())
421 .await
422 {
423 Ok(id) => Ok(id),
424 Err(e) if Self::is_transport_error(&e.to_string()) => {
425 let restarted = self.restart_bridge().await?;
426 self.register_with_channel(restarted, schema_id, xsd_bytes)
427 .await
428 }
429 Err(e) => Err(e),
430 }
431 }
432
433 async fn validate(&self, schema_id: &str, doc_bytes: Vec<u8>) -> Result<(), ValidatorError> {
434 let channel = self.ensure_bridge_ready().await?;
435 let req = ValidateWithRequest {
436 schema_id: schema_id.to_string(),
437 document: doc_bytes.clone(),
438 };
439
440 let response = match self.rpc.validate_with(channel.clone(), req).await {
441 Ok(resp) => resp,
442 Err(e) if Self::is_transport_error(&e.to_string()) => {
443 let restarted = self.restart_bridge().await?;
444 self.rpc
445 .validate_with(
446 restarted,
447 ValidateWithRequest {
448 schema_id: schema_id.to_string(),
449 document: doc_bytes,
450 },
451 )
452 .await?
453 }
454 Err(e) => return Err(e),
455 };
456
457 if let Some(err) = response.error {
458 return Err(ValidatorError::from_bridge_error(&err));
459 }
460 if response.valid {
461 return Ok(());
462 }
463
464 let details = response
465 .errors
466 .iter()
467 .map(|e| format!("{}:{} {}", e.line, e.column, e.message))
468 .collect::<Vec<_>>()
469 .join("\n");
470 Err(ValidatorError::validation(format!(
471 "XSD validation failed:\n{details}"
472 )))
473 }
474}
475
476impl Default for XsdBridgeBackend {
477 fn default() -> Self {
478 Self::new()
479 }
480}
481
482#[cfg(test)]
483mod tests {
484 use super::*;
485 use std::sync::atomic::{AtomicUsize, Ordering};
486 use tonic::transport::Endpoint;
487
488 #[derive(Debug)]
489 struct MockRpc {
490 register_calls: Arc<AtomicUsize>,
491 validate_ok: bool,
492 }
493
494 #[async_trait]
495 impl XsdBridgeRpc for MockRpc {
496 async fn register_schema(
497 &self,
498 _channel: Channel,
499 request: RegisterSchemaRequest,
500 ) -> Result<RegisterSchemaResponse, ValidatorError> {
501 self.register_calls.fetch_add(1, Ordering::SeqCst);
502 Ok(RegisterSchemaResponse {
503 schema_id: request.schema_id,
504 error: None,
505 })
506 }
507
508 async fn validate_with(
509 &self,
510 _channel: Channel,
511 _request: ValidateWithRequest,
512 ) -> Result<ValidateResponse, ValidatorError> {
513 Ok(ValidateResponse {
514 valid: self.validate_ok,
515 errors: Vec::new(),
516 error: None,
517 })
518 }
519 }
520
521 fn lazy_channel() -> Channel {
522 Endpoint::from_static("http://127.0.0.1:65535").connect_lazy()
523 }
524
525 #[tokio::test]
526 async fn xsd_bridge_reconnect_reseeds() {
527 let calls = Arc::new(AtomicUsize::new(0));
528 let rpc = Arc::new(MockRpc {
529 register_calls: Arc::clone(&calls),
530 validate_ok: true,
531 });
532 let channel = lazy_channel();
533 let backend = XsdBridgeBackend::for_test(rpc, channel);
534
535 let _id_a = backend.register(b"<xsd:a/>".to_vec()).await.unwrap();
536 let _id_b = backend.register(b"<xsd:b/>".to_vec()).await.unwrap();
537
538 backend.on_reconnect(&lazy_channel()).unwrap();
539 tokio::time::sleep(Duration::from_millis(20)).await;
540
541 assert!(calls.load(Ordering::SeqCst) >= 4);
542 }
543}