1use std::{fmt, sync::Arc};
2
3use async_lock::RwLock;
4use axum::{
5 Router,
6 body::Bytes,
7 extract::State,
8 http::{HeaderMap, StatusCode},
9 response::Response,
10 routing::{get, post},
11};
12use hotshot_types::{
13 light_client::{
14 LCV1StateSignatureRequestBody, LCV1StateSignaturesBundle, LCV2StateSignatureRequestBody,
15 LCV2StateSignaturesBundle, LCV3StateSignatureRequestBody, LCV3StateSignaturesBundle,
16 },
17 traits::signature_key::LCV1StateSignatureKey,
18};
19use http_wire::{self as wire, DecodeFailure, WireError, cors_layer, healthcheck_response};
20use lcv1_relay::{LCV1StateRelayServerDataSource, LCV1StateRelayServerState};
21use lcv2_relay::{LCV2StateRelayServerDataSource, LCV2StateRelayServerState};
22use lcv3_relay::{LCV3StateRelayServerDataSource, LCV3StateRelayServerState};
23use serde::{Deserialize, Serialize, de::DeserializeOwned};
24use tokio::{net::TcpListener, sync::oneshot};
25use url::Url;
26use vbs::version::StaticVersionType;
27
28pub mod lcv1_relay;
29pub mod lcv2_relay;
30pub mod lcv3_relay;
31pub mod stake_table_tracker;
32
33#[derive(Debug, Clone, Serialize, Deserialize)]
37pub struct RelayError {
38 pub status: u16,
39 pub message: String,
40}
41
42impl RelayError {
43 pub fn catch_all(status: StatusCode, message: impl Into<String>) -> Self {
44 Self {
45 status: status.as_u16(),
46 message: message.into(),
47 }
48 }
49}
50
51impl fmt::Display for RelayError {
52 fn fmt(&self, f: &mut fmt::Formatter<'_>) -> fmt::Result {
53 write!(f, "Error {}: {}", self.status, self.message)
54 }
55}
56
57impl std::error::Error for RelayError {}
58
59impl WireError for RelayError {
60 fn status(&self) -> StatusCode {
61 StatusCode::from_u16(self.status).unwrap_or(StatusCode::INTERNAL_SERVER_ERROR)
62 }
63
64 fn catch_all(status: StatusCode, message: String) -> Self {
65 RelayError::catch_all(status, message)
66 }
67}
68
69fn decode_body<T: DeserializeOwned>(headers: &HeaderMap, body: &[u8]) -> Result<T, RelayError> {
70 wire::decode_body(headers, body).map_err(|failure| match failure {
71 DecodeFailure::Binary(e) => {
72 RelayError::catch_all(StatusCode::BAD_REQUEST, format!("invalid binary body: {e}"))
73 },
74 DecodeFailure::Json(e) => {
75 RelayError::catch_all(StatusCode::BAD_REQUEST, format!("invalid json body: {e}"))
76 },
77 DecodeFailure::UnsupportedContentType => RelayError::catch_all(
78 StatusCode::BAD_REQUEST,
79 "missing or unsupported Content-Type",
80 ),
81 })
82}
83
84pub struct StateRelayServerState {
86 lcv1_state: LCV1StateRelayServerState,
88 lcv2_state: LCV2StateRelayServerState,
90 lcv3_state: LCV3StateRelayServerState,
92 shutdown: Option<oneshot::Receiver<()>>,
94}
95
96impl StateRelayServerState {
97 pub fn new(sequencer_url: Url) -> Self {
99 let stake_table_tracker =
100 Arc::new(stake_table_tracker::StakeTableTracker::new(sequencer_url));
101 Self {
102 lcv1_state: LCV1StateRelayServerState::new(stake_table_tracker.clone()),
103 lcv2_state: LCV2StateRelayServerState::new(stake_table_tracker.clone()),
104 lcv3_state: LCV3StateRelayServerState::new(stake_table_tracker),
105 shutdown: None,
106 }
107 }
108
109 pub fn with_shutdown_signal(
110 mut self,
111 shutdown_listener: Option<oneshot::Receiver<()>>,
112 ) -> Self {
113 if self.shutdown.is_some() {
114 panic!("A shutdown signal is already registered and can not be registered twice");
115 }
116 self.shutdown = shutdown_listener;
117 self
118 }
119}
120
121#[async_trait::async_trait]
122impl LCV1StateRelayServerDataSource for StateRelayServerState {
123 fn get_latest_signature_bundle(&self) -> Result<LCV1StateSignaturesBundle, RelayError> {
124 self.lcv1_state.get_latest_signature_bundle()
125 }
126
127 async fn post_signature(
128 &mut self,
129 req: LCV1StateSignatureRequestBody,
130 ) -> Result<(), RelayError> {
131 self.lcv1_state.post_signature(req).await
132 }
133}
134
135#[async_trait::async_trait]
136impl LCV2StateRelayServerDataSource for StateRelayServerState {
137 fn get_latest_signature_bundle(&self) -> Result<LCV2StateSignaturesBundle, RelayError> {
138 self.lcv2_state.get_latest_signature_bundle()
139 }
140
141 async fn post_signature(
142 &mut self,
143 req: LCV2StateSignatureRequestBody,
144 ) -> Result<(), RelayError> {
145 self.lcv2_state.post_signature(req).await
146 }
147}
148
149#[async_trait::async_trait]
150impl LCV3StateRelayServerDataSource for StateRelayServerState {
151 fn get_latest_signature_bundle(&self) -> Result<LCV3StateSignaturesBundle, RelayError> {
152 self.lcv3_state.get_latest_signature_bundle()
153 }
154
155 async fn post_signature(
156 &mut self,
157 req: LCV3StateSignatureRequestBody,
158 ) -> Result<(), RelayError> {
159 self.lcv3_state.post_signature(req).await
160 }
161}
162
163type SharedState = Arc<RwLock<StateRelayServerState>>;
165
166async fn post_state_signature(
169 state: &SharedState,
170 headers: &HeaderMap,
171 body: &[u8],
172) -> Result<(), RelayError> {
173 if let Ok(req) = decode_body::<LCV3StateSignatureRequestBody>(headers, body) {
174 tracing::debug!("Received LCV3 state signature: {req}");
175 let mut state = state.write().await;
176 if let Err(e) =
177 LCV2StateRelayServerDataSource::post_signature(&mut *state, req.clone().into()).await
178 {
179 tracing::error!("Failed to post downgraded LCV2 state signature: {}", e);
180 }
181 LCV3StateRelayServerDataSource::post_signature(&mut *state, req).await
182 } else if let Ok(req) = decode_body::<LCV2StateSignatureRequestBody>(headers, body) {
183 tracing::debug!("Received LCV2 state signature: {req}");
184 let mut state = state.write().await;
185 if LCV1StateSignatureKey::verify_state_sig(&req.key, &req.signature, &req.state) {
186 LCV1StateRelayServerDataSource::post_signature(&mut *state, req.into()).await
187 } else {
188 LCV2StateRelayServerDataSource::post_signature(&mut *state, req).await
189 }
190 } else if let Ok(req) = decode_body::<LCV1StateSignatureRequestBody>(headers, body) {
191 tracing::debug!("Received LCV1 state signature: {req}");
192 let mut state = state.write().await;
193 LCV1StateRelayServerDataSource::post_signature(&mut *state, req).await
194 } else {
195 Err(RelayError::catch_all(
196 StatusCode::BAD_REQUEST,
197 "Invalid request body",
198 ))
199 }
200}
201
202async fn post_legacy_state_signature(
205 state: &SharedState,
206 headers: &HeaderMap,
207 body: &[u8],
208) -> Result<(), RelayError> {
209 let req = if let Ok(req) = decode_body::<LCV1StateSignatureRequestBody>(headers, body) {
210 req
211 } else if let Ok(req) = decode_body::<LCV2StateSignatureRequestBody>(headers, body) {
212 req.into()
213 } else {
214 return Err(RelayError::catch_all(
215 StatusCode::BAD_REQUEST,
216 "Invalid request body",
217 ));
218 };
219 let mut state = state.write().await;
220 LCV1StateRelayServerDataSource::post_signature(&mut *state, req).await
221}
222
223async fn get_latest_state(state: &SharedState) -> Result<LCV2StateSignaturesBundle, RelayError> {
226 let state = state.read().await;
227 LCV2StateRelayServerDataSource::get_latest_signature_bundle(&*state)
228}
229
230async fn get_latest_legacy_state(
232 state: &SharedState,
233) -> Result<LCV2StateSignaturesBundle, RelayError> {
234 let state = state.read().await;
235 LCV1StateRelayServerDataSource::get_latest_signature_bundle(&*state)
236 .map(LCV2StateSignaturesBundle::from_v1)
237}
238
239async fn get_latest_state_v1(state: &SharedState) -> Result<LCV1StateSignaturesBundle, RelayError> {
240 let state = state.read().await;
241 LCV1StateRelayServerDataSource::get_latest_signature_bundle(&*state)
242}
243
244async fn get_latest_state_v2(state: &SharedState) -> Result<LCV2StateSignaturesBundle, RelayError> {
245 let state = state.read().await;
246 LCV2StateRelayServerDataSource::get_latest_signature_bundle(&*state)
247}
248
249async fn get_latest_state_v3(state: &SharedState) -> Result<LCV3StateSignaturesBundle, RelayError> {
250 let state = state.read().await;
251 LCV3StateRelayServerDataSource::get_latest_signature_bundle(&*state)
252}
253
254async fn healthcheck(headers: HeaderMap) -> Response {
255 healthcheck_response(&headers)
256}
257
258async fn post_state(State(state): State<SharedState>, headers: HeaderMap, body: Bytes) -> Response {
259 wire::respond(
260 &headers,
261 post_state_signature(&state, &headers, &body).await,
262 )
263}
264
265async fn get_state(State(state): State<SharedState>, headers: HeaderMap) -> Response {
266 wire::respond(&headers, get_latest_state(&state).await)
267}
268
269async fn post_legacy_state(
270 State(state): State<SharedState>,
271 headers: HeaderMap,
272 body: Bytes,
273) -> Response {
274 wire::respond(
275 &headers,
276 post_legacy_state_signature(&state, &headers, &body).await,
277 )
278}
279
280async fn get_legacy_state(State(state): State<SharedState>, headers: HeaderMap) -> Response {
281 wire::respond(&headers, get_latest_legacy_state(&state).await)
282}
283
284async fn get_lateststate_v1(State(state): State<SharedState>, headers: HeaderMap) -> Response {
285 wire::respond(&headers, get_latest_state_v1(&state).await)
286}
287
288async fn get_lateststate_v2(State(state): State<SharedState>, headers: HeaderMap) -> Response {
289 wire::respond(&headers, get_latest_state_v2(&state).await)
290}
291
292async fn get_lateststate_v3(State(state): State<SharedState>, headers: HeaderMap) -> Response {
293 wire::respond(&headers, get_latest_state_v3(&state).await)
294}
295
296const STATE_PATH: &str = "/api/state";
297const LEGACY_STATE_PATH: &str = "/api/legacy-state";
298const LATEST_STATE_PATH: &str = "/api/lateststate";
299
300fn router(state: SharedState) -> Router {
306 let mut router = Router::<SharedState>::new()
307 .route("/healthcheck", get(healthcheck))
308 .route(STATE_PATH, post(post_state).get(get_state))
309 .route(
310 LEGACY_STATE_PATH,
311 post(post_legacy_state).get(get_legacy_state),
312 )
313 .route(LATEST_STATE_PATH, get(get_lateststate_v3))
314 .route(&format!("/v1{LATEST_STATE_PATH}"), get(get_lateststate_v1))
315 .route(&format!("/v2{LATEST_STATE_PATH}"), get(get_lateststate_v2))
316 .route(&format!("/v3{LATEST_STATE_PATH}"), get(get_lateststate_v3));
317 for v in 1..=3 {
318 router = router
319 .route(
320 &format!("/v{v}{STATE_PATH}"),
321 post(post_state).get(get_state),
322 )
323 .route(
324 &format!("/v{v}{LEGACY_STATE_PATH}"),
325 post(post_legacy_state).get(get_legacy_state),
326 );
327 }
328 router.with_state(state).layer(cors_layer())
329}
330
331async fn serve(server_url: Url, state: StateRelayServerState) -> anyhow::Result<()> {
332 let host = server_url
333 .host_str()
334 .ok_or_else(|| anyhow::anyhow!("relay server url missing host: {server_url}"))?;
335 let port = server_url
336 .port_or_known_default()
337 .ok_or_else(|| anyhow::anyhow!("relay server url missing port: {server_url}"))?;
338 let listener = TcpListener::bind((host, port)).await?;
339
340 tracing::info!(%server_url, "Relay server starts serving at ");
341 axum::serve(listener, router(Arc::new(RwLock::new(state)))).await?;
342 Ok(())
343}
344
345pub async fn run_relay_server<BindVer: StaticVersionType + 'static>(
346 shutdown_listener: Option<oneshot::Receiver<()>>,
347 sequencer_url: Url,
348 url: Url,
349 _bind_version: BindVer,
352) -> anyhow::Result<()> {
353 let state = StateRelayServerState::new(sequencer_url).with_shutdown_signal(shutdown_listener);
354 serve(url, state).await
355}
356
357pub async fn run_relay_server_with_state<BindVer: StaticVersionType + 'static>(
358 server_url: Url,
359 _bind_version: BindVer,
360 state: StateRelayServerState,
361) -> anyhow::Result<()> {
362 serve(server_url, state).await
363}
364
365#[cfg(test)]
366mod test {
367 use alloy::primitives::{FixedBytes, U256};
368 use axum::Json;
369 use espresso_types::SeqTypes;
370 use hotshot::types::SchnorrPubKey;
371 use hotshot_contract_adapter::light_client::derive_signed_state_digest;
372 use hotshot_types::{
373 PeerConfig, ValidatorConfig,
374 light_client::{LightClientState, StakeTableState},
375 traits::signature_key::{LCV2StateSignatureKey, LCV3StateSignatureKey},
376 };
377 use http_client::{Client, error::ClientErr};
378 use vbs::version::StaticVersion;
379
380 use super::*;
381
382 type TestApiVer = StaticVersion<0, 1>;
383
384 async fn spawn_fake_sequencer(peer: PeerConfig<SeqTypes>) -> Url {
388 let config = serde_json::json!({
389 "config": {
390 "known_nodes_with_stake": [peer],
391 "epoch_height": 0,
392 "epoch_start_block": 0,
393 }
394 });
395 let app = Router::new().route(
396 "/config/hotshot",
397 get(move || {
398 let config = config.clone();
399 async move { Json(config) }
400 }),
401 );
402 let listener = TcpListener::bind("127.0.0.1:0").await.unwrap();
403 let addr = listener.local_addr().unwrap();
404 tokio::spawn(async move {
405 axum::serve(listener, app).await.unwrap();
406 });
407 format!("http://{addr}").parse().unwrap()
408 }
409
410 async fn spawn_relay(state: StateRelayServerState) -> Url {
412 let listener = TcpListener::bind("127.0.0.1:0").await.unwrap();
413 let addr = listener.local_addr().unwrap();
414 tokio::spawn(async move {
415 axum::serve(listener, router(Arc::new(RwLock::new(state))))
416 .await
417 .unwrap();
418 });
419 format!("http://{addr}").parse().unwrap()
420 }
421
422 #[tokio::test]
426 async fn post_and_fetch_lcv3_state_signature() {
427 let validator = ValidatorConfig::<SeqTypes>::generated_from_seed_indexed(
428 [7u8; 32],
429 0,
430 U256::from(1),
431 true,
432 );
433 let sequencer_url = spawn_fake_sequencer(validator.public_config()).await;
434 let relay_url = spawn_relay(StateRelayServerState::new(sequencer_url)).await;
435
436 let light_client_state = LightClientState {
437 view_number: 1,
438 block_height: 1,
439 block_comm_root: Default::default(),
440 };
441 let next_stake = StakeTableState::default();
442 let auth_root = FixedBytes::<32>::default();
443 let digest = derive_signed_state_digest(&light_client_state, &next_stake, &auth_root);
444 let signature = <SchnorrPubKey as LCV3StateSignatureKey>::sign_state(
445 &validator.state_private_key,
446 digest,
447 )
448 .unwrap();
449 let v2_signature = <SchnorrPubKey as LCV2StateSignatureKey>::sign_state(
450 &validator.state_private_key,
451 &light_client_state,
452 &next_stake,
453 )
454 .unwrap();
455 let request_body = LCV3StateSignatureRequestBody {
456 key: validator.state_public_key.clone(),
457 state: light_client_state,
458 next_stake,
459 auth_root,
460 signature,
461 v2_signature,
462 };
463
464 let client = Client::<ClientErr, TestApiVer>::new(relay_url);
465 client
466 .post::<()>("api/state")
467 .body_binary(&request_body)
468 .unwrap()
469 .send()
470 .await
471 .unwrap();
472
473 let bundle = client
474 .get::<LCV3StateSignaturesBundle>("api/lateststate")
475 .send()
476 .await
477 .unwrap();
478 assert_eq!(bundle.state, light_client_state);
479 assert_eq!(bundle.signatures.len(), 1);
480 assert!(bundle.signatures.contains_key(&validator.state_public_key));
481 }
482
483 #[tokio::test]
486 async fn responses_carry_cors_headers() {
487 let relay_url = spawn_relay(StateRelayServerState::new(
488 "http://127.0.0.1:1".parse().unwrap(),
489 ))
490 .await;
491 for path in ["healthcheck", "api/lateststate", "no/such/route"] {
492 let resp = reqwest::Client::new()
493 .get(relay_url.join(path).unwrap())
494 .header("Origin", "https://example.com")
495 .send()
496 .await
497 .unwrap();
498 assert_eq!(
499 resp.headers()
500 .get("access-control-allow-origin")
501 .unwrap_or_else(|| panic!("no CORS header on {path}")),
502 "*",
503 "{path}"
504 );
505 }
506 }
507}