1313// See the License for the specific language governing permissions and
1414// limitations under the License.
1515
16+ use std:: num:: NonZeroUsize ;
1617use std:: path:: Path ;
1718use std:: str:: FromStr ;
1819use std:: sync:: Arc ;
@@ -21,12 +22,13 @@ use std::time::{Duration, Instant};
2122use anyhow:: { Context , ensure} ;
2223use futures:: StreamExt ;
2324use serde:: { Deserialize , Serialize } ;
25+ use serde_json:: Value ;
26+ use sse_core:: { SseDecoder , SseEvent } ;
2427use tokio:: sync:: Semaphore ;
2528
2629use crate :: manifest:: { Manifest , ManifestRequest } ;
2730
28- mod sse;
29- use sse:: CompletionStream ;
31+ const MAX_SSE_PAYLOAD_BYTES : usize = 1024 * 1024 ;
3032
3133#[ derive( Debug , Clone ) ]
3234pub struct DriveConfig {
@@ -221,30 +223,84 @@ async fn execute_request(
221223 ) ;
222224 }
223225 let mut stream = response. bytes_stream ( ) ;
224- let mut output = CompletionStream :: default ( ) ;
226+ let mut decoder = SseDecoder :: with_limit ( NonZeroUsize :: new ( MAX_SSE_PAYLOAD_BYTES ) . unwrap ( ) ) ;
225227 while let Some ( chunk) = stream. next ( ) . await {
226- let bytes = match chunk {
228+ let mut bytes = match chunk {
227229 Ok ( bytes) => bytes,
228230 Err ( error) => return result. finish ( dispatch_time, Some ( error. to_string ( ) ) ) ,
229231 } ;
230- let progress = output. push ( & bytes) ;
231- result. observed_output_tokens = output. output_tokens ( ) ;
232- match progress {
233- Ok ( true ) if result. first_output_ms . is_none ( ) => {
234- result. first_output_ms = Some ( duration_ms ( dispatch_time. elapsed ( ) ) ) ;
232+ while let Some ( event) = decoder. next ( & mut bytes) {
233+ let event = match event {
234+ Ok ( SseEvent :: Message ( event) ) => event,
235+ Ok ( SseEvent :: Retry ( _) ) => continue ,
236+ Err ( error) => return result. finish ( dispatch_time, Some ( error. to_string ( ) ) ) ,
237+ } ;
238+ match record_completion_event ( & mut result, & event. data , dispatch_time) {
239+ Ok ( true ) => {
240+ result. ok = true ;
241+ return result. finish ( dispatch_time, None ) ;
242+ }
243+ Ok ( false ) => { }
244+ Err ( error) => return result. finish ( dispatch_time, Some ( error. to_string ( ) ) ) ,
235245 }
236- Ok ( _) => { }
237- Err ( error) => return result. finish ( dispatch_time, Some ( error. to_string ( ) ) ) ,
238246 }
239247 }
240- if !output. is_complete ( ) {
241- return result. finish (
242- dispatch_time,
243- Some ( "upstream SSE response ended before [DONE]" . into ( ) ) ,
248+ result. finish (
249+ dispatch_time,
250+ Some ( "upstream SSE response ended before [DONE]" . into ( ) ) ,
251+ )
252+ }
253+
254+ fn record_completion_event (
255+ result : & mut RequestResult ,
256+ data : & str ,
257+ dispatch_time : Instant ,
258+ ) -> anyhow:: Result < bool > {
259+ let data = data. trim ( ) ;
260+ if data == "[DONE]" {
261+ return Ok ( true ) ;
262+ }
263+ if data. is_empty ( ) {
264+ return Ok ( false ) ;
265+ }
266+ let value: Value = serde_json:: from_str ( data) . context ( "invalid upstream SSE JSON" ) ?;
267+ ensure ! (
268+ value. get( "error" ) . is_none_or( Value :: is_null) ,
269+ "upstream returned an SSE error event"
270+ ) ;
271+ let generated_output = value[ "choices" ] . as_array ( ) . is_some_and ( |choices| {
272+ choices. iter ( ) . any ( |choice| {
273+ let delta = & choice[ "delta" ] ;
274+ [ "content" , "reasoning_content" , "reasoning" ]
275+ . iter ( )
276+ . any ( |field| delta[ * field] . as_str ( ) . is_some_and ( |text| !text. is_empty ( ) ) )
277+ || delta[ "tool_calls" ] . as_array ( ) . is_some_and ( |calls| {
278+ calls. iter ( ) . any ( |call| {
279+ call[ "function" ] [ "arguments" ]
280+ . as_str ( )
281+ . is_some_and ( |arguments| !arguments. is_empty ( ) )
282+ } )
283+ } )
284+ } )
285+ } ) ;
286+ if generated_output && result. first_output_ms . is_none ( ) {
287+ result. first_output_ms = Some ( duration_ms ( dispatch_time. elapsed ( ) ) ) ;
288+ }
289+ if let Some ( tokens) = value
290+ . pointer ( "/usage/completion_tokens" )
291+ . or_else ( || value. get ( "output_tokens_so_far" ) )
292+ . filter ( |tokens| !tokens. is_null ( ) )
293+ {
294+ result. observed_output_tokens = Some (
295+ tokens
296+ . as_u64 ( )
297+ . context ( "upstream output token usage is not an unsigned integer" ) ?,
244298 ) ;
299+ } else if generated_output {
300+ // A prior cumulative counter does not cover later uncounted output.
301+ result. observed_output_tokens = None ;
245302 }
246- result. ok = true ;
247- result. finish ( dispatch_time, None )
303+ Ok ( false )
248304}
249305
250306fn duration_ms ( duration : Duration ) -> u64 {
@@ -325,12 +381,13 @@ mod tests {
325381 }
326382 }
327383
328- async fn drive_test_response ( body : & ' static str ) -> ( RequestResult , serde_json:: Value ) {
384+ async fn drive_test_response ( body : & str ) -> ( RequestResult , serde_json:: Value ) {
329385 let listener = TcpListener :: bind ( "127.0.0.1:0" ) . await . unwrap ( ) ;
330386 let endpoint = format ! (
331387 "http://{}/v1/chat/completions" ,
332388 listener. local_addr( ) . unwrap( )
333389 ) ;
390+ let body = body. to_owned ( ) ;
334391 let server = tokio:: spawn ( async move {
335392 let ( mut socket, _) = listener. accept ( ) . await . unwrap ( ) ;
336393 let request = read_test_request ( & mut socket) . await ;
@@ -362,7 +419,7 @@ mod tests {
362419
363420 #[ tokio:: test]
364421 async fn complete_response_records_actual_usage_and_requests_usage_reporting ( ) {
365- let ( result, request) = drive_test_response ( "data: {\" choices\" :[{\" delta\" :{\" content\" :\" answer \ " }}], \ " usage\" :{\" completion_tokens\" :2}}\n \ n data: [DONE]\n \n " ) . await ;
422+ let ( result, request) = drive_test_response ( "\u{feff} data: {\" choices\" :[{\" delta\" :{\" content\" :\" \u{03bb} \ " }}]} \r \n \r \n data: { \ " usage\" :{\" completion_tokens\" :2}}\r \n \r \ n data: [DONE]\r \n \r \n " ) . await ;
366423 assert ! ( result. ok, "{:?}" , result. error) ;
367424 assert_eq ! ( result. output_tokens, 100 ) ;
368425 assert_eq ! ( result. observed_output_tokens, Some ( 2 ) ) ;
@@ -384,6 +441,37 @@ mod tests {
384441 }
385442 }
386443
444+ #[ tokio:: test]
445+ async fn role_only_completion_keeps_output_timing_and_usage_unknown ( ) {
446+ let ( result, _) = drive_test_response ( ": keepalive\n \n data: {\" choices\" :[{\" delta\" :{\" role\" :\" assistant\" ,\" content\" :\" \" }}]}\n \n data: [DONE]\n \n " ) . await ;
447+ assert ! ( result. ok, "{:?}" , result. error) ;
448+ assert_eq ! ( result. first_output_ms, None ) ;
449+ assert_eq ! ( result. observed_output_tokens, None ) ;
450+ }
451+
452+ #[ tokio:: test]
453+ async fn cumulative_usage_must_cover_the_last_generated_output ( ) {
454+ let ( result, _) = drive_test_response ( "data: {\" choices\" :[{\" delta\" :{\" content\" :\" one\" }}],\" output_tokens_so_far\" :1}\n \n data: {\" choices\" :[{\" delta\" :{\" content\" :\" two\" }}]}\n \n data: [DONE]\n \n " ) . await ;
455+ assert ! ( result. ok, "{:?}" , result. error) ;
456+ assert ! ( result. first_output_ms. is_some( ) ) ;
457+ assert_eq ! ( result. observed_output_tokens, None ) ;
458+ }
459+
460+ #[ tokio:: test]
461+ async fn invalid_or_oversized_events_fail_the_request ( ) {
462+ for data in [
463+ "{bad}" . to_owned ( ) ,
464+ r#"{"error":{"message":"failed"}}"# . to_owned ( ) ,
465+ r#"{"usage":{"completion_tokens":-1}}"# . to_owned ( ) ,
466+ serde_json:: json!( { "choices" : [ { "delta" : { "content" : "x" . repeat( MAX_SSE_PAYLOAD_BYTES ) } } ] } ) . to_string ( ) ,
467+ ] {
468+ let ( result, _) =
469+ drive_test_response ( & format ! ( "data: {data}\n \n data: [DONE]\n \n " ) ) . await ;
470+ assert ! ( !result. ok) ;
471+ assert ! ( result. error. is_some( ) ) ;
472+ }
473+ }
474+
387475 #[ tokio:: test]
388476 async fn scheduled_sleep_does_not_hold_concurrency_permit ( ) {
389477 let listener = TcpListener :: bind ( "127.0.0.1:0" )
0 commit comments