@@ -3,7 +3,12 @@ use std::{
33 convert:: Infallible ,
44 fmt:: { self , Debug } ,
55 net:: SocketAddr ,
6- sync:: { Arc , Mutex } ,
6+ pin:: Pin ,
7+ sync:: {
8+ atomic:: { AtomicU64 , Ordering } ,
9+ Arc , Mutex ,
10+ } ,
11+ task:: { Context , Poll } ,
712} ;
813
914use axum:: {
@@ -26,7 +31,10 @@ use serde_json::{json, Value};
2631use tokio:: sync:: { OwnedSemaphorePermit , Semaphore } ;
2732use url:: Url ;
2833
29- use crate :: { broadcast, gafe} ;
34+ use http_body:: Frame ;
35+ use pin_project_lite:: pin_project;
36+
37+ use crate :: { broadcast, gafe, user_query} ;
3038
3139macro_rules! user_error {
3240 ( $e: expr) => {
@@ -127,6 +135,7 @@ pub enum Error {
127135 User ( String ) ,
128136 Timeout ( Option < String > ) ,
129137 TooManyRequests ( Option < String > ) ,
138+ TransferExceeded ( Option < String > ) ,
130139
131140 Server ( Box < dyn std:: error:: Error + Send + Sync > ) ,
132141}
@@ -150,6 +159,10 @@ impl Serialize for Error {
150159 state. serialize_field ( "error" , "too_many_requests" ) ?;
151160 state. serialize_field ( "message" , & opt_msg) ?;
152161 }
162+ Error :: TransferExceeded ( opt_msg) => {
163+ state. serialize_field ( "error" , "transfer_exceeded" ) ?;
164+ state. serialize_field ( "message" , & opt_msg) ?;
165+ }
153166 Error :: Server ( err) => {
154167 state. serialize_field ( "error" , "server" ) ?;
155168 state. serialize_field ( "message" , & err. to_string ( ) ) ?;
@@ -167,6 +180,8 @@ impl std::fmt::Display for Error {
167180 Error :: Timeout ( None ) => write ! ( f, "Operation timed out" ) ,
168181 Error :: TooManyRequests ( Some ( msg) ) => write ! ( f, "Too many requests: {msg}" ) ,
169182 Error :: TooManyRequests ( None ) => write ! ( f, "Too many requests" ) ,
183+ Error :: TransferExceeded ( Some ( msg) ) => write ! ( f, "Transfer exceeded: {msg}" ) ,
184+ Error :: TransferExceeded ( None ) => write ! ( f, "Transfer limit exceeded" ) ,
170185 Error :: Server ( err) => write ! ( f, "Server error: {err}" ) ,
171186 }
172187 }
@@ -219,6 +234,12 @@ impl axum::response::IntoResponse for Error {
219234 StatusCode :: TOO_MANY_REQUESTS ,
220235 msg. unwrap_or ( String :: from ( "too many requests" ) ) ,
221236 ) ,
237+ Self :: TransferExceeded ( msg) => (
238+ StatusCode :: TOO_MANY_REQUESTS ,
239+ msg. unwrap_or ( String :: from (
240+ "Transfer limit exceeded. Upgrade at: https://www.indexsupply.net" ,
241+ ) ) ,
242+ ) ,
222243 Self :: User ( msg) => ( StatusCode :: BAD_REQUEST , msg) ,
223244 Self :: Server ( e) => {
224245 tracing:: error!( %e, "server-error={:?}" , e) ;
@@ -285,10 +306,21 @@ pub async fn limit(
285306 "Rate limited. Create or upgrade API Key at: https://www.indexsupply.net" ,
286307 ) ) ) ) ;
287308 }
288- match tokio:: time:: timeout ( account_limit. timeout , next. run ( request) ) . await {
289- Ok ( response) => Ok ( response) ,
290- Err ( _) => Err ( Error :: Timeout ( None ) ) ,
291- }
309+ let log = request. extensions ( ) . get :: < user_query:: RequestLog > ( ) . cloned ( ) ;
310+ let response = match tokio:: time:: timeout ( account_limit. timeout , next. run ( request) ) . await {
311+ Ok ( response) => response,
312+ Err ( _) => return Err ( Error :: Timeout ( None ) ) ,
313+ } ;
314+ let ( parts, body) = response. into_parts ( ) ;
315+ let counting = CountingBody {
316+ inner : body,
317+ count : Arc :: new ( AtomicU64 :: new ( 0 ) ) ,
318+ log,
319+ } ;
320+ Ok ( axum:: http:: Response :: from_parts (
321+ parts,
322+ axum:: body:: Body :: new ( counting) ,
323+ ) )
292324}
293325
294326#[ derive( Clone , Copy , Default , Debug , Deserialize , Eq , Hash , PartialEq , Serialize ) ]
@@ -496,19 +528,45 @@ pub async fn latency_header(
496528 Ok ( response)
497529}
498530
499- pub async fn content_length_header (
500- request : axum:: extract:: Request ,
501- next : axum:: middleware:: Next ,
502- ) -> Result < axum:: response:: Response , Error > {
503- let response = next. run ( request) . await ;
504- let span = tracing:: Span :: current ( ) ;
505- response
506- . headers ( )
507- . get ( "content-length" )
508- . and_then ( |cl| cl. to_str ( ) . ok ( ) )
509- . map ( |cl| cl. parse :: < u64 > ( ) . ok ( ) )
510- . map ( |size| span. record ( "size" , size) ) ;
511- Ok ( response)
531+ pin_project ! {
532+ pub struct CountingBody <B > {
533+ #[ pin]
534+ inner: B ,
535+ count: Arc <AtomicU64 >,
536+ log: Option <user_query:: RequestLog >,
537+ }
538+ }
539+
540+ impl < B > http_body:: Body for CountingBody < B >
541+ where
542+ B : http_body:: Body < Data = bytes:: Bytes > ,
543+ {
544+ type Data = B :: Data ;
545+ type Error = B :: Error ;
546+
547+ fn poll_frame (
548+ self : Pin < & mut Self > ,
549+ cx : & mut Context < ' _ > ,
550+ ) -> Poll < Option < Result < Frame < Self :: Data > , Self :: Error > > > {
551+ let this = self . project ( ) ;
552+ match this. inner . poll_frame ( cx) {
553+ Poll :: Ready ( Some ( Ok ( frame) ) ) => {
554+ if let Some ( data) = frame. data_ref ( ) {
555+ this. count . fetch_add ( data. len ( ) as u64 , Ordering :: Relaxed ) ;
556+ }
557+ Poll :: Ready ( Some ( Ok ( frame) ) )
558+ }
559+ Poll :: Ready ( None ) => {
560+ let total = this. count . load ( Ordering :: Relaxed ) ;
561+ tracing:: Span :: current ( ) . record ( "size" , total) ;
562+ if let Some ( log) = this. log . take ( ) {
563+ log. set_bytes ( total) ;
564+ }
565+ Poll :: Ready ( None )
566+ }
567+ other => other,
568+ }
569+ }
512570}
513571
514572pub async fn log_fields (
0 commit comments