Skip to main content

sol_parser_sdk/shredstream/
client.rs

1//! ShredStream 客户端
2//!
3//! `solana_entry::entry::Entry` 在 Agave SDK 中带 `deprecated`(需显式启用不稳定 feature 才消除);
4//! 本模块仍依赖其 bincode 布局解码 Shred 侧 `entries` 负载。
5#![allow(deprecated)]
6
7use std::sync::atomic::{AtomicU64, Ordering};
8use std::sync::Arc;
9use std::time::Duration;
10
11use crossbeam_queue::ArrayQueue;
12use futures::StreamExt;
13use solana_entry::entry::Entry as SolanaEntry;
14use solana_sdk::message::VersionedMessage;
15use tokio::sync::Mutex;
16use tokio::task::JoinHandle;
17use tonic::transport::{Channel, Endpoint};
18
19use crate::core::now_micros;
20use crate::grpc::types::EventTypeFilter;
21use crate::shredstream::config::ShredStreamConfig;
22use crate::shredstream::proto::{Entry, ShredstreamProxyClient, SubscribeEntriesRequest};
23use crate::DexEvent;
24
25static SHREDSTREAM_DROPPED_EVENTS: AtomicU64 = AtomicU64::new(0);
26
27enum EventSink<'a> {
28    Queue(&'a Arc<ArrayQueue<DexEvent>>),
29    Callback(&'a (dyn Fn(DexEvent) + Send + Sync)),
30}
31
32impl EventSink<'_> {
33    #[inline]
34    fn deliver(&self, event: DexEvent) {
35        match self {
36            EventSink::Queue(queue) => {
37                if queue.push(event).is_err() {
38                    record_shredstream_dropped_event();
39                }
40            }
41            EventSink::Callback(callback) => callback(event),
42        }
43    }
44}
45
46#[inline]
47fn record_shredstream_dropped_event() -> u64 {
48    let dropped = SHREDSTREAM_DROPPED_EVENTS.fetch_add(1, Ordering::Relaxed) + 1;
49    if dropped <= 10 || dropped.is_power_of_two() {
50        log::warn!(
51            target: "sol_parser_sdk::shredstream",
52            "ShredStream event queue is full; dropped event count={}",
53            dropped
54        );
55    }
56    dropped
57}
58
59/// ShredStream 客户端
60#[derive(Clone)]
61pub struct ShredStreamClient {
62    endpoint: String,
63    config: ShredStreamConfig,
64    subscription_handle: Arc<Mutex<Option<JoinHandle<()>>>>,
65    subscription_lifecycle: Arc<Mutex<()>>,
66}
67
68impl ShredStreamClient {
69    /// 创建新客户端
70    pub async fn new(endpoint: impl Into<String>) -> crate::common::AnyResult<Self> {
71        Self::new_with_config(endpoint, ShredStreamConfig::default()).await
72    }
73
74    /// 使用自定义配置创建客户端
75    pub async fn new_with_config(
76        endpoint: impl Into<String>,
77        config: ShredStreamConfig,
78    ) -> crate::common::AnyResult<Self> {
79        let endpoint = endpoint.into();
80        // 测试连接
81        let _ = Self::connect_client(&endpoint, &config).await?;
82
83        Ok(Self {
84            endpoint,
85            config,
86            subscription_handle: Arc::new(Mutex::new(None)),
87            subscription_lifecycle: Arc::new(Mutex::new(())),
88        })
89    }
90
91    /// 订阅 DEX 事件(自动重连)
92    ///
93    /// 返回一个队列,事件会被推送到该队列中
94    pub async fn subscribe(&self) -> crate::common::AnyResult<Arc<ArrayQueue<DexEvent>>> {
95        self.subscribe_with_filter(None).await
96    }
97
98    /// 订阅 DEX 事件,并在 ShredStream 热路径中按 SDK 事件类型提前过滤。
99    ///
100    /// 过滤发生在解析分发前,用于低延迟场景避免解析不需要的协议/事件。
101    pub async fn subscribe_with_filter(
102        &self,
103        event_type_filter: Option<EventTypeFilter>,
104    ) -> crate::common::AnyResult<Arc<ArrayQueue<DexEvent>>> {
105        let _lifecycle = self.subscription_lifecycle.lock().await;
106        self.stop_without_lifecycle_lock().await;
107
108        let queue = Arc::new(ArrayQueue::new(100_000));
109        let queue_clone = Arc::clone(&queue);
110
111        let endpoint = self.endpoint.clone();
112        let config = self.config.clone();
113
114        let handle = tokio::spawn(async move {
115            let mut delay = config.reconnect_delay_ms;
116            let mut attempts = 0u32;
117
118            loop {
119                if config.max_reconnect_attempts > 0 && attempts >= config.max_reconnect_attempts {
120                    log::error!("Max reconnection attempts reached, giving up");
121                    break;
122                }
123                attempts += 1;
124
125                match Self::stream_events(
126                    &endpoint,
127                    &config,
128                    event_type_filter.as_ref(),
129                    EventSink::Queue(&queue_clone),
130                )
131                .await
132                {
133                    Ok(_) => {
134                        log::warn!("ShredStream ended cleanly - reconnecting in {}ms", delay);
135                        tokio::time::sleep(tokio::time::Duration::from_millis(delay.max(1))).await;
136                        delay = config.reconnect_delay_ms;
137                        attempts = 0;
138                    }
139                    Err(e) => {
140                        log::error!("ShredStream error: {} - retry in {}ms", e, delay);
141                        tokio::time::sleep(tokio::time::Duration::from_millis(delay)).await;
142                        delay = (delay * 2).min(60_000);
143                    }
144                }
145            }
146        });
147
148        *self.subscription_handle.lock().await = Some(handle);
149        Ok(queue)
150    }
151
152    /// 订阅 DEX 事件,并在解析热路径中直接回调事件,避免跨任务队列调度。
153    ///
154    /// 这是最低延迟路径;回调会在 ShredStream 读流任务内执行,应避免阻塞 I/O 或重计算。
155    pub async fn subscribe_with_filter_callback<F>(
156        &self,
157        event_type_filter: Option<EventTypeFilter>,
158        callback: F,
159    ) -> crate::common::AnyResult<()>
160    where
161        F: Fn(DexEvent) + Send + Sync + 'static,
162    {
163        let _lifecycle = self.subscription_lifecycle.lock().await;
164        self.stop_without_lifecycle_lock().await;
165
166        let endpoint = self.endpoint.clone();
167        let config = self.config.clone();
168        let callback = Arc::new(callback);
169
170        let handle = tokio::spawn(async move {
171            let mut delay = config.reconnect_delay_ms;
172            let mut attempts = 0u32;
173
174            loop {
175                if config.max_reconnect_attempts > 0 && attempts >= config.max_reconnect_attempts {
176                    log::error!("Max reconnection attempts reached, giving up");
177                    break;
178                }
179                attempts += 1;
180
181                match Self::stream_events_callback(
182                    &endpoint,
183                    &config,
184                    event_type_filter.as_ref(),
185                    callback.clone(),
186                )
187                .await
188                {
189                    Ok(_) => {
190                        log::warn!("ShredStream ended cleanly - reconnecting in {}ms", delay);
191                        tokio::time::sleep(tokio::time::Duration::from_millis(delay.max(1))).await;
192                        delay = config.reconnect_delay_ms;
193                        attempts = 0;
194                    }
195                    Err(e) => {
196                        log::error!("ShredStream error: {} - retry in {}ms", e, delay);
197                        tokio::time::sleep(tokio::time::Duration::from_millis(delay)).await;
198                        delay = (delay * 2).min(60_000);
199                    }
200                }
201            }
202        });
203
204        *self.subscription_handle.lock().await = Some(handle);
205        Ok(())
206    }
207
208    /// 停止订阅
209    pub async fn stop(&self) {
210        let _lifecycle = self.subscription_lifecycle.lock().await;
211        self.stop_without_lifecycle_lock().await;
212    }
213
214    async fn stop_without_lifecycle_lock(&self) {
215        if let Some(handle) = self.subscription_handle.lock().await.take() {
216            handle.abort();
217            let _ = handle.await;
218        }
219    }
220
221    async fn connect_client(
222        endpoint: &str,
223        config: &ShredStreamConfig,
224    ) -> crate::common::AnyResult<ShredstreamProxyClient<Channel>> {
225        let mut builder = Endpoint::from_shared(endpoint.to_string())?;
226        if config.connection_timeout_ms > 0 {
227            builder = builder.connect_timeout(Duration::from_millis(config.connection_timeout_ms));
228        }
229        let channel = builder.connect().await?;
230        Ok(ShredstreamProxyClient::new(channel)
231            .max_decoding_message_size(config.max_decoding_message_size))
232    }
233
234    /// 核心事件流处理
235    async fn stream_events(
236        endpoint: &str,
237        config: &ShredStreamConfig,
238        event_type_filter: Option<&EventTypeFilter>,
239        sink: EventSink<'_>,
240    ) -> Result<(), String> {
241        let mut client = Self::connect_client(endpoint, config).await.map_err(|e| e.to_string())?;
242        let request = tonic::Request::new(SubscribeEntriesRequest {});
243        let response = if config.request_timeout_ms > 0 {
244            tokio::time::timeout(
245                Duration::from_millis(config.request_timeout_ms),
246                client.subscribe_entries(request),
247            )
248            .await
249            .map_err(|_| {
250                format!(
251                    "ShredStream subscribe request timed out after {}ms",
252                    config.request_timeout_ms
253                )
254            })?
255            .map_err(|e| e.to_string())?
256        } else {
257            client.subscribe_entries(request).await.map_err(|e| e.to_string())?
258        };
259        let mut stream = response.into_inner();
260
261        log::info!("ShredStream connected, receiving entries...");
262
263        let mut events = Vec::with_capacity(4);
264        while let Some(message) = stream.next().await {
265            match message {
266                Ok(entry) => {
267                    Self::process_entry(entry, event_type_filter, &sink, &mut events);
268                }
269                Err(e) => {
270                    log::error!("Stream error: {:?}", e);
271                    return Err(e.to_string());
272                }
273            }
274        }
275
276        Ok(())
277    }
278
279    async fn stream_events_callback(
280        endpoint: &str,
281        config: &ShredStreamConfig,
282        event_type_filter: Option<&EventTypeFilter>,
283        callback: Arc<dyn Fn(DexEvent) + Send + Sync>,
284    ) -> Result<(), String> {
285        Self::stream_events(
286            endpoint,
287            config,
288            event_type_filter,
289            EventSink::Callback(callback.as_ref()),
290        )
291        .await
292    }
293
294    /// 处理单个 Entry 消息
295    #[inline]
296    fn process_entry(
297        entry: Entry,
298        event_type_filter: Option<&EventTypeFilter>,
299        sink: &EventSink<'_>,
300        events: &mut Vec<DexEvent>,
301    ) {
302        let slot = entry.slot;
303        let recv_us = now_micros();
304
305        // 反序列化 Entry 数据
306        let entries = match wincode::deserialize::<Vec<SolanaEntry>>(&entry.entries) {
307            Ok(e) => e,
308            Err(e) => {
309                log::debug!("Failed to deserialize entries: {}", e);
310                return;
311            }
312        };
313
314        // 处理每个 Entry 中的交易
315        let mut tx_index = 0u64;
316        for entry in entries {
317            for transaction in entry.transactions.iter() {
318                if transaction.sanitize().is_err() {
319                    log::debug!("Ignoring unsanitized ShredStream transaction");
320                    tx_index += 1;
321                    continue;
322                }
323                events.clear();
324                Self::process_transaction(
325                    transaction,
326                    slot,
327                    recv_us,
328                    tx_index,
329                    event_type_filter,
330                    events,
331                    sink,
332                );
333                tx_index += 1;
334            }
335        }
336    }
337
338    /// 处理单个交易
339    #[inline]
340    fn process_transaction(
341        transaction: &solana_sdk::transaction::VersionedTransaction,
342        slot: u64,
343        recv_us: i64,
344        tx_index: u64,
345        event_type_filter: Option<&EventTypeFilter>,
346        events: &mut Vec<DexEvent>,
347        sink: &EventSink<'_>,
348    ) {
349        Self::parse_transaction_events(
350            transaction,
351            slot,
352            recv_us,
353            tx_index,
354            event_type_filter,
355            events,
356        );
357
358        for event in events.drain(..) {
359            sink.deliver(event);
360        }
361    }
362
363    #[inline]
364    fn parse_transaction_events(
365        transaction: &solana_sdk::transaction::VersionedTransaction,
366        slot: u64,
367        recv_us: i64,
368        tx_index: u64,
369        event_type_filter: Option<&EventTypeFilter>,
370        events: &mut Vec<DexEvent>,
371    ) {
372        let Some(&signature) = transaction.signatures.first() else {
373            return;
374        };
375        if let VersionedMessage::V0(m) = &transaction.message {
376            if !m.address_table_lookups.is_empty() {
377                log::trace!(
378                    target: "sol_parser_sdk::shredstream",
379                    "V0 tx uses address lookup tables; shred parser will use static accounts and default placeholders for ALT-loaded accounts"
380                );
381            }
382        }
383        // 热路径:`static_account_keys` 零拷贝、`pump_ix` 内不克隆 CompiledInstruction。
384        super::pump_ix::parse_transaction_dex_events_with_filter(
385            transaction,
386            signature,
387            slot,
388            tx_index,
389            recv_us,
390            event_type_filter,
391            events,
392        );
393        crate::core::pumpfun_fee_enrich::enrich_pumpfun_same_tx_post_merge(events);
394
395        for event in events.iter_mut() {
396            if let Some(meta) = event.metadata_mut() {
397                meta.grpc_recv_us = recv_us;
398            }
399        }
400    }
401}
402
403#[cfg(test)]
404mod tests {
405    use super::*;
406    use crate::core::events::{EventMetadata, PumpFunCreateTokenEvent};
407    use crate::instr::program_ids::PUMPFUN_PROGRAM_ID;
408    use solana_sdk::hash::Hash;
409    use solana_sdk::message::{
410        compiled_instruction::CompiledInstruction, v0, MessageHeader, VersionedMessage,
411    };
412    use solana_sdk::pubkey::Pubkey;
413    use solana_sdk::signature::Signature;
414    use solana_sdk::transaction::VersionedTransaction;
415    use std::sync::atomic::{AtomicUsize, Ordering as AtomicOrdering};
416
417    fn push_string(data: &mut Vec<u8>, value: &str) {
418        data.extend_from_slice(&(value.len() as u32).to_le_bytes());
419        data.extend_from_slice(value.as_bytes());
420    }
421
422    fn pumpfun_create_data() -> Vec<u8> {
423        let mut data = Vec::new();
424        data.extend_from_slice(&[24, 30, 200, 40, 5, 28, 7, 119]);
425        push_string(&mut data, "Callback Test");
426        push_string(&mut data, "CBT");
427        push_string(&mut data, "https://example.invalid/callback.json");
428        data.extend_from_slice(Pubkey::new_unique().as_ref());
429        data
430    }
431
432    fn pumpfun_create_tx() -> VersionedTransaction {
433        let mut account_keys = (0..10).map(|_| Pubkey::new_unique()).collect::<Vec<_>>();
434        account_keys.push(PUMPFUN_PROGRAM_ID);
435
436        VersionedTransaction {
437            signatures: vec![Signature::default()],
438            message: VersionedMessage::V0(v0::Message {
439                header: MessageHeader {
440                    num_required_signatures: 1,
441                    num_readonly_signed_accounts: 0,
442                    num_readonly_unsigned_accounts: 1,
443                },
444                account_keys,
445                recent_blockhash: Hash::default(),
446                instructions: vec![CompiledInstruction::new_from_raw_parts(
447                    10,
448                    pumpfun_create_data(),
449                    (0..10).collect(),
450                )],
451                address_table_lookups: Vec::new(),
452            }),
453        }
454    }
455
456    #[test]
457    fn dropped_counter_increments_without_panicking() {
458        let before = SHREDSTREAM_DROPPED_EVENTS.load(Ordering::Relaxed);
459        let queue = ArrayQueue::new(1);
460
461        queue
462            .push(DexEvent::PumpFunCreate(PumpFunCreateTokenEvent {
463                metadata: EventMetadata::default(),
464                ..Default::default()
465            }))
466            .expect("first push fits");
467
468        if queue
469            .push(DexEvent::PumpFunCreate(PumpFunCreateTokenEvent {
470                metadata: EventMetadata::default(),
471                ..Default::default()
472            }))
473            .is_err()
474        {
475            record_shredstream_dropped_event();
476        }
477
478        assert!(SHREDSTREAM_DROPPED_EVENTS.load(Ordering::Relaxed) > before);
479    }
480
481    #[test]
482    fn callback_path_delivers_events_without_queue() {
483        let entries = vec![SolanaEntry {
484            num_hashes: 1,
485            hash: Hash::default(),
486            transactions: vec![pumpfun_create_tx()],
487        }];
488        let entry = Entry { slot: 42, entries: wincode::serialize(&entries).unwrap() };
489        let count = AtomicUsize::new(0);
490
491        let mut events = Vec::with_capacity(4);
492        ShredStreamClient::process_entry(
493            entry,
494            None,
495            &EventSink::Callback(&|event| {
496                assert!(matches!(event, DexEvent::PumpFunCreate(_)));
497                assert_eq!(event.metadata().slot, 42);
498                count.fetch_add(1, AtomicOrdering::Relaxed);
499            }),
500            &mut events,
501        );
502
503        assert_eq!(count.load(AtomicOrdering::Relaxed), 1);
504    }
505
506    #[test]
507    fn wincode_decodes_solana4_v1_entries() {
508        let entries = vec![SolanaEntry {
509            num_hashes: 1,
510            hash: Hash::new_unique(),
511            transactions: vec![VersionedTransaction {
512                signatures: vec![Signature::from([9; 64])],
513                message: VersionedMessage::V1(solana_sdk::message::v1::Message {
514                    header: MessageHeader { num_required_signatures: 1, ..Default::default() },
515                    config: solana_sdk::message::v1::TransactionConfig::empty()
516                        .with_compute_unit_limit(250_000)
517                        .with_priority_fee(1_500),
518                    lifetime_specifier: Hash::new_unique(),
519                    account_keys: vec![Pubkey::new_unique()],
520                    instructions: Vec::new(),
521                }),
522            }],
523        }];
524        let bytes = wincode::serialize(&entries).expect("serialize V1 Entry");
525        let decoded: Vec<SolanaEntry> = wincode::deserialize(&bytes).expect("decode V1 Entry");
526
527        assert_eq!(decoded, entries);
528    }
529
530    #[test]
531    fn process_entry_rejects_unsanitized_transactions() {
532        let mut transaction = pumpfun_create_tx();
533        let VersionedMessage::V0(message) = &mut transaction.message else {
534            panic!("expected V0 fixture");
535        };
536        message.instructions[0].program_id_index = u8::MAX;
537        let entries = vec![SolanaEntry {
538            num_hashes: 1,
539            hash: Hash::default(),
540            transactions: vec![transaction],
541        }];
542        let entry = Entry { slot: 42, entries: wincode::serialize(&entries).unwrap() };
543        let count = AtomicUsize::new(0);
544        let mut events = Vec::with_capacity(4);
545
546        ShredStreamClient::process_entry(
547            entry,
548            None,
549            &EventSink::Callback(&|_| {
550                count.fetch_add(1, AtomicOrdering::Relaxed);
551            }),
552            &mut events,
553        );
554
555        assert_eq!(count.load(AtomicOrdering::Relaxed), 0);
556    }
557}