use std::{
sync::Arc,
time::{SystemTime, UNIX_EPOCH},
};
use dynamo_protocols::types::{ChatCompletionStreamOptions, CompletionUsage};
use crate::protocols::common::{
extensions::{NvExt, NvExtResponseFieldSelection},
timing::RequestTracker,
};
#[derive(Debug, Clone, Default)]
pub struct DeltaGeneratorOptions {
pub enable_usage: bool,
pub continuous_usage_stats: bool,
pub enable_logprobs: bool,
pub return_tokens_as_token_ids: bool,
pub response_fields: NvExtResponseFieldSelection,
}
impl DeltaGeneratorOptions {
pub fn new(
stream_options: Option<&ChatCompletionStreamOptions>,
return_tokens_as_token_ids: Option<bool>,
enable_logprobs: bool,
nvext: Option<&NvExt>,
) -> Self {
let response_fields = NvExtResponseFieldSelection::from_nvext(nvext);
DeltaGeneratorOptions {
enable_usage: stream_options.is_some_and(|opts| opts.include_usage),
continuous_usage_stats: stream_options.is_some_and(|opts| opts.continuous_usage_stats),
enable_logprobs,
response_fields,
return_tokens_as_token_ids: return_tokens_as_token_ids.unwrap_or(false),
}
}
}
pub(crate) fn initial_state() -> (u32, CompletionUsage, Arc<RequestTracker>) {
let now_time = SystemTime::now()
.duration_since(UNIX_EPOCH)
.unwrap() .as_secs();
let now: u32 = now_time.try_into().expect("timestamp exceeds u32::MAX");
let usage = dynamo_protocols::types::CompletionUsage {
completion_tokens: 0,
prompt_tokens: 0,
total_tokens: 0,
completion_tokens_details: None,
prompt_tokens_details: None,
};
let tracker = Arc::new(RequestTracker::new());
(now, usage, tracker)
}
pub(crate) fn enable_usage_for_nonstreaming(
stream_options: &mut Option<ChatCompletionStreamOptions>,
original_stream_flag: bool,
) {
if original_stream_flag {
return;
}
stream_options
.get_or_insert_with(|| ChatCompletionStreamOptions {
include_usage: true,
continuous_usage_stats: false,
})
.include_usage = true;
}
pub(crate) fn force_include_usage(stream_options: &mut Option<ChatCompletionStreamOptions>) {
stream_options
.get_or_insert_with(|| ChatCompletionStreamOptions {
include_usage: true,
continuous_usage_stats: false,
})
.include_usage = true;
}
#[cfg(test)]
mod tests {
use super::*;
#[test]
fn force_include_usage_inserts_missing_options() {
let mut options = None;
force_include_usage(&mut options);
let options = options.expect("stream options should be inserted");
assert!(options.include_usage);
assert!(!options.continuous_usage_stats);
}
#[test]
fn force_include_usage_overrides_false_and_preserves_siblings() {
let mut options = Some(ChatCompletionStreamOptions {
include_usage: false,
continuous_usage_stats: true,
});
force_include_usage(&mut options);
let options = options.expect("stream options should remain present");
assert!(options.include_usage);
assert!(options.continuous_usage_stats);
}
}