Skip to main content

relay_event_normalization/normalize/span/
ai.rs

1//! AI cost calculation.
2
3use crate::eap::AttributesLike;
4use crate::statsd::{Counters, map_origin_to_integration, platform_tag};
5use crate::{ModelCostV2, ModelMetadata};
6use relay_conventions::attributes::*;
7use relay_event_schema::protocol::{
8    Event, Measurements, OperationType, Span, SpanData, TraceContext,
9};
10use relay_protocol::{Annotated, Getter, Value};
11
12/// Amount of used tokens for a model call.
13#[derive(Debug, Copy, Clone)]
14pub struct UsedTokens {
15    /// Total amount of input tokens used.
16    pub input_tokens: f64,
17    /// Amount of cached tokens used.
18    ///
19    /// This is a subset of [`Self::input_tokens`].
20    pub input_cached_tokens: f64,
21    /// Amount of cache write tokens used.
22    ///
23    /// This is a subset of [`Self::input_tokens`].
24    pub input_cache_write_tokens: f64,
25    /// Total amount of output tokens.
26    pub output_tokens: f64,
27    /// Total amount of reasoning tokens.
28    ///
29    /// This is a subset of [`Self::output_tokens`].
30    pub output_reasoning_tokens: f64,
31}
32
33impl UsedTokens {
34    /// Extracts [`UsedTokens`] from [`SpanData`] attributes.
35    pub fn from_span_data(data: &SpanData) -> Self {
36        macro_rules! get_value {
37            ($e:expr) => {
38                data.get_value($e).and_then(|v| v.as_f64()).unwrap_or(0.0)
39            };
40        }
41
42        Self {
43            input_tokens: get_value!(GEN_AI__USAGE__INPUT_TOKENS),
44            output_tokens: get_value!(GEN_AI__USAGE__OUTPUT_TOKENS),
45            output_reasoning_tokens: get_value!(GEN_AI__USAGE__REASONING__OUTPUT_TOKENS),
46            input_cached_tokens: get_value!(GEN_AI__USAGE__CACHE_READ__INPUT_TOKENS),
47            input_cache_write_tokens: get_value!(GEN_AI__USAGE__CACHE_CREATION__INPUT_TOKENS),
48        }
49    }
50
51    /// Returns `true` if any tokens were used.
52    pub fn has_usage(&self) -> bool {
53        self.input_tokens > 0.0 || self.output_tokens > 0.0
54    }
55
56    /// Calculates the total amount of input tokens billed at the standard rate.
57    ///
58    /// Both [`Self::input_cached_tokens`] and [`Self::input_cache_write_tokens`] are
59    /// subsets of [`Self::input_tokens`] and are billed separately at their own
60    /// (cached / cache-write) rates, so both are subtracted here to avoid charging
61    /// them twice.
62    pub fn raw_input_tokens(&self) -> f64 {
63        self.input_tokens - self.input_cached_tokens - self.input_cache_write_tokens
64    }
65
66    /// Calculates the total amount of raw, non-reasoning output tokens.
67    ///
68    /// Subtracts reasoning tokens from the total token count.
69    pub fn raw_output_tokens(&self) -> f64 {
70        self.output_tokens - self.output_reasoning_tokens
71    }
72}
73
74/// Calculated model call costs.
75#[derive(Debug, Copy, Clone)]
76pub struct CalculatedCost {
77    /// The total cost of all input tokens (raw + cached + cache_write).
78    pub input: f64,
79    /// The total cost of all output tokens (raw + reasoning).
80    pub output: f64,
81    /// The cost of cached input tokens only (subset of `input`).
82    pub cache_read_input: f64,
83    /// The cost of cache-write input tokens only (subset of `input`).
84    pub cache_creation_input: f64,
85    /// The cost of reasoning output tokens only (subset of `output`).
86    pub reasoning_output: f64,
87}
88
89impl CalculatedCost {
90    /// The total, input and output, cost.
91    pub fn total(&self) -> f64 {
92        self.input + self.output
93    }
94}
95
96/// Calculates the total cost for a model call.
97///
98/// Returns `None` if no tokens were used.
99pub fn calculate_costs(
100    model_cost: &ModelCostV2,
101    tokens: UsedTokens,
102    integration: &str,
103    platform: &str,
104) -> Option<CalculatedCost> {
105    if !tokens.has_usage() {
106        relay_statsd::metric!(
107            counter(Counters::GenAiCostCalculationResult) += 1,
108            result = "calculation_no_tokens",
109            integration = integration,
110            platform = platform,
111        );
112        return None;
113    }
114
115    let cache_read_input = tokens.input_cached_tokens * model_cost.input_cached_per_token;
116    let cache_creation_input =
117        tokens.input_cache_write_tokens * model_cost.input_cache_write_per_token;
118    let input = (tokens.raw_input_tokens() * model_cost.input_per_token)
119        + cache_read_input
120        + cache_creation_input;
121
122    // For now most of the models do not differentiate between reasoning and output token cost,
123    // it costs the same.
124    let reasoning_per_token = match model_cost.output_reasoning_per_token {
125        r if r > 0.0 => r,
126        _ => model_cost.output_per_token,
127    };
128    let reasoning_output = tokens.output_reasoning_tokens * reasoning_per_token;
129    let output = (tokens.raw_output_tokens() * model_cost.output_per_token) + reasoning_output;
130
131    let metric_label = match (input, output) {
132        (x, y) if x < 0.0 || y < 0.0 => "calculation_negative",
133        (0.0, 0.0) => "calculation_zero",
134        _ => "calculation_positive",
135    };
136
137    relay_statsd::metric!(
138        counter(Counters::GenAiCostCalculationResult) += 1,
139        result = metric_label,
140        integration = integration,
141        platform = platform,
142    );
143
144    Some(CalculatedCost {
145        input,
146        output,
147        cache_read_input,
148        cache_creation_input,
149        reasoning_output,
150    })
151}
152
153/// Default AI operation stored in [`GEN_AI__OPERATION__TYPE`]
154/// for AI spans without a well known AI span op.
155///
156/// See also: [`infer_ai_operation_type`].
157pub const DEFAULT_AI_OPERATION: &str = "ai_client";
158
159/// Infers the AI operation from an AI operation name.
160///
161/// The operation name is usually inferred from the
162/// [`GEN_AI__OPERATION__NAME`] span attribute and the span
163/// operation.
164///
165/// Sentry expects the operation type in the [`GEN_AI__OPERATION__TYPE`] attribute.
166///
167/// The function returns `None` when the op is not a well known AI operation, callers likely want to default
168/// the value to [`DEFAULT_AI_OPERATION`] for AI spans.
169pub fn infer_ai_operation_type(op_name: &str) -> Option<&'static str> {
170    let ai_op = match op_name {
171        // Full matches:
172        "ai.run.generateText"
173        | "ai.run.generateObject"
174        | "gen_ai.invoke_agent"
175        | "ai.pipeline.generate_text"
176        | "ai.pipeline.generate_object"
177        | "ai.pipeline.stream_text"
178        | "ai.pipeline.stream_object"
179        | "gen_ai.create_agent"
180        | "invoke_agent"
181        | "create_agent" => "agent",
182        "gen_ai.execute_tool" | "execute_tool" => "tool",
183        "gen_ai.handoff" | "handoff" => "handoff",
184        "ai.processor" | "processor_run" => "other",
185        // Prefix matches:
186        op if op.starts_with("ai.streamText.doStream") => "ai_client",
187        op if op.starts_with("ai.streamText") => "agent",
188
189        op if op.starts_with("ai.generateText.doGenerate") => "ai_client",
190        op if op.starts_with("ai.generateText") => "agent",
191
192        op if op.starts_with("ai.generateObject.doGenerate") => "ai_client",
193        op if op.starts_with("ai.generateObject") => "agent",
194
195        op if op.starts_with("ai.toolCall") => "tool",
196        // No match:
197        _ => return None,
198    };
199
200    Some(ai_op)
201}
202
203/// Returns whether a valid total cost is attached.
204pub fn has_valid_total_cost(attributes: &impl AttributesLike) -> bool {
205    attributes
206        .get_value(GEN_AI__COST__TOTAL_TOKENS)
207        .and_then(Value::as_f64)
208        .is_some()
209}
210
211/// Calculates the cost of an AI model based on the model cost and the tokens used.
212/// Calculated cost is in US dollars.
213fn extract_ai_model_cost_data(
214    model_cost: Option<&ModelCostV2>,
215    data: &mut SpanData,
216    origin: Option<&str>,
217    platform: Option<&str>,
218) {
219    // Preserve existing total cost instead of recalculating and overwriting it.
220    if has_valid_total_cost(data) {
221        return;
222    }
223
224    let integration = map_origin_to_integration(origin);
225    let platform = platform_tag(platform);
226
227    let Some(model_cost) = model_cost else {
228        relay_statsd::metric!(
229            counter(Counters::GenAiCostCalculationResult) += 1,
230            result = "calculation_no_model_cost_available",
231            integration = integration,
232            platform = platform,
233        );
234        return;
235    };
236
237    let used_tokens = UsedTokens::from_span_data(&*data);
238    let Some(costs) = calculate_costs(model_cost, used_tokens, integration, platform) else {
239        return;
240    };
241
242    data.other
243        .entry(GEN_AI__COST__TOTAL_TOKENS.to_owned())
244        .or_default()
245        .set_value(Value::F64(costs.total()).into());
246
247    // Set individual cost components
248    data.other
249        .entry(GEN_AI__COST__INPUT_TOKENS.to_owned())
250        .or_default()
251        .set_value(Value::F64(costs.input).into());
252    data.other
253        .entry(GEN_AI__COST__CACHE_READ__INPUT_TOKENS.to_owned())
254        .or_default()
255        .set_value(Value::F64(costs.cache_read_input).into());
256    data.other
257        .entry(GEN_AI__COST__CACHE_CREATION__INPUT_TOKENS.to_owned())
258        .or_default()
259        .set_value(Value::F64(costs.cache_creation_input).into());
260
261    data.other
262        .entry(GEN_AI__COST__OUTPUT_TOKENS.to_owned())
263        .or_default()
264        .set_value(Value::F64(costs.output).into());
265
266    data.other
267        .entry(GEN_AI__COST__REASONING__OUTPUT_TOKENS.to_owned())
268        .or_default()
269        .set_value(Value::F64(costs.reasoning_output).into());
270}
271
272/// Maps AI-related measurements (legacy) to span data.
273fn map_ai_measurements_to_data(data: &mut SpanData, measurements: Option<&Measurements>) {
274    let set_field_from_measurement = |target_field: &mut Annotated<Value>,
275                                      measurement_key: &str| {
276        if let Some(measurements) = measurements
277            && target_field.value().is_none()
278            && let Some(value) = measurements.get_value(measurement_key)
279        {
280            target_field.set_value(Value::F64(value.to_f64()).into());
281        }
282    };
283
284    set_field_from_measurement(
285        data.other
286            .entry(GEN_AI__USAGE__TOTAL_TOKENS.to_owned())
287            .or_default(),
288        "ai_total_tokens_used",
289    );
290    set_field_from_measurement(
291        data.other
292            .entry(GEN_AI__USAGE__INPUT_TOKENS.to_owned())
293            .or_default(),
294        "ai_prompt_tokens_used",
295    );
296    set_field_from_measurement(
297        data.other
298            .entry(GEN_AI__USAGE__OUTPUT_TOKENS.to_owned())
299            .or_default(),
300        "ai_completion_tokens_used",
301    );
302}
303
304fn set_total_tokens(data: &mut SpanData) {
305    // It might be that 'total_tokens' is not set in which case we need to calculate it
306    if data.get_value(GEN_AI__USAGE__TOTAL_TOKENS).is_none() {
307        let input_tokens = data
308            .get_value(GEN_AI__USAGE__INPUT_TOKENS)
309            .and_then(Value::as_f64);
310        let output_tokens = data
311            .get_value(GEN_AI__USAGE__OUTPUT_TOKENS)
312            .and_then(Value::as_f64);
313
314        if input_tokens.is_none() && output_tokens.is_none() {
315            // don't set total_tokens if there are no input nor output tokens
316            return;
317        }
318
319        data.other
320            .entry(GEN_AI__USAGE__TOTAL_TOKENS.to_owned())
321            .or_default()
322            .set_value(
323                Value::F64(input_tokens.unwrap_or(0.0) + output_tokens.unwrap_or(0.0)).into(),
324            );
325    }
326}
327
328/// Sets the context window size and utilization for the model.
329fn extract_context_utilization(data: &mut SpanData, model_metadata: &ModelMetadata) {
330    let model_id = data.get_str(GEN_AI__RESPONSE__MODEL);
331
332    let context_size = model_id.and_then(|id| model_metadata.context_size(id));
333
334    let Some(context_size) = context_size else {
335        return;
336    };
337
338    data.other
339        .entry(GEN_AI__CONTEXT__WINDOW_SIZE.to_owned())
340        .or_default()
341        .set_value(Value::U64(context_size).into());
342
343    let total_tokens = data
344        .get_value(GEN_AI__USAGE__TOTAL_TOKENS)
345        .and_then(Value::as_f64);
346
347    if let Some(total_tokens) = total_tokens {
348        data.other
349            .entry(GEN_AI__CONTEXT__UTILIZATION.to_owned())
350            .or_default()
351            .set_value(Value::F64(total_tokens / context_size as f64).into());
352    }
353}
354
355/// Extract the additional data into the span
356fn extract_ai_data(
357    data: &mut SpanData,
358    duration: f64,
359    model_metadata: &ModelMetadata,
360    origin: Option<&str>,
361    platform: Option<&str>,
362) {
363    // Extracts the response tokens per second
364    if data
365        .get_value(GEN_AI__RESPONSE__TOKENS_PER_SECOND)
366        .is_none()
367        && duration > 0.0
368        && let Some(output_tokens) = data
369            .get_value(GEN_AI__USAGE__OUTPUT_TOKENS)
370            .and_then(Value::as_f64)
371    {
372        data.other
373            .entry(GEN_AI__RESPONSE__TOKENS_PER_SECOND.to_owned())
374            .or_default()
375            .set_value(Value::F64(output_tokens / (duration / 1000.0)).into());
376    }
377
378    extract_context_utilization(data, model_metadata);
379
380    // Extracts the total cost of the AI model used
381    if let Some(model_id) = data.get_str(GEN_AI__RESPONSE__MODEL) {
382        extract_ai_model_cost_data(
383            model_metadata.cost_per_token(model_id),
384            data,
385            origin,
386            platform,
387        )
388    } else {
389        relay_statsd::metric!(
390            counter(Counters::GenAiCostCalculationResult) += 1,
391            result = "calculation_no_model_id_available",
392            integration = map_origin_to_integration(origin),
393            platform = platform_tag(platform),
394        );
395    }
396}
397
398/// Enrich the AI span data
399fn enrich_ai_span_data(
400    span_data: &mut Annotated<SpanData>,
401    span_op: &Annotated<OperationType>,
402    measurements: &Annotated<Measurements>,
403    duration: f64,
404    model_metadata: Option<&ModelMetadata>,
405    origin: Option<&str>,
406    platform: Option<&str>,
407) {
408    if !is_ai_span(span_data, span_op.value()) {
409        return;
410    }
411
412    let data = span_data.get_or_insert_with(SpanData::default);
413
414    map_ai_measurements_to_data(data, measurements.value());
415
416    set_total_tokens(data);
417
418    // Default response model to request model if not set.
419    if data.get_value(GEN_AI__RESPONSE__MODEL).is_none()
420        && let Some(request_model) = data.get_value(GEN_AI__REQUEST__MODEL).cloned()
421    {
422        data.other
423            .entry(GEN_AI__RESPONSE__MODEL.to_owned())
424            .or_default()
425            .set_value(Some(request_model));
426    }
427
428    // Default agent name to function_id if not set.
429    if data.get_value(GEN_AI__AGENT__NAME).is_none()
430        && let Some(function_id) = data.get_value(GEN_AI__FUNCTION_ID).cloned()
431    {
432        data.other
433            .entry(GEN_AI__AGENT__NAME.to_owned())
434            .or_default()
435            .set_value(Some(function_id));
436    }
437
438    if let Some(model_metadata) = model_metadata {
439        extract_ai_data(data, duration, model_metadata, origin, platform);
440    } else {
441        relay_statsd::metric!(
442            counter(Counters::GenAiCostCalculationResult) += 1,
443            result = "calculation_no_model_cost_available",
444            integration = map_origin_to_integration(origin),
445            platform = platform_tag(platform),
446        );
447    }
448
449    let ai_op_type = data
450        .get_str(GEN_AI__OPERATION__NAME)
451        .or(span_op.value().map(String::as_str))
452        .and_then(infer_ai_operation_type)
453        .unwrap_or(DEFAULT_AI_OPERATION);
454
455    data.other
456        .entry(GEN_AI__OPERATION__TYPE.to_owned())
457        .or_default()
458        .set_value(Some(Value::String(ai_op_type.to_owned())));
459}
460
461/// Enrich the AI span data
462pub fn enrich_ai_span(span: &mut Span, model_metadata: Option<&ModelMetadata>) {
463    let duration = span
464        .get_value("span.duration")
465        .and_then(|v| v.as_f64())
466        .unwrap_or(0.0);
467
468    enrich_ai_span_data(
469        &mut span.data,
470        &span.op,
471        &span.measurements,
472        duration,
473        model_metadata,
474        span.origin.as_str(),
475        span.platform.as_str(),
476    );
477}
478
479/// Extract the ai data from all of an event's spans
480pub fn enrich_ai_event_data(event: &mut Event, model_metadata: Option<&ModelMetadata>) {
481    let event_duration = event
482        .get_value("event.duration")
483        .and_then(|v| v.as_f64())
484        .unwrap_or(0.0);
485
486    if let Some(trace_context) = event
487        .contexts
488        .value_mut()
489        .as_mut()
490        .and_then(|c| c.get_mut::<TraceContext>())
491    {
492        enrich_ai_span_data(
493            &mut trace_context.data,
494            &trace_context.op,
495            &event.measurements,
496            event_duration,
497            model_metadata,
498            trace_context.origin.as_str(),
499            event.platform.as_str(),
500        );
501    }
502    let spans = event.spans.value_mut().iter_mut().flatten();
503    let spans = spans.filter_map(|span| span.value_mut().as_mut());
504
505    for span in spans {
506        let span_duration = span
507            .get_value("span.duration")
508            .and_then(|v| v.as_f64())
509            .unwrap_or(0.0);
510        let span_platform = span.platform.as_str().or_else(|| event.platform.as_str());
511
512        enrich_ai_span_data(
513            &mut span.data,
514            &span.op,
515            &span.measurements,
516            span_duration,
517            model_metadata,
518            span.origin.as_str(),
519            span_platform,
520        );
521    }
522}
523
524/// Returns true if the span is an AI span.
525/// AI spans are spans with either a gen_ai.operation.name attribute or op starting with "ai."
526/// (legacy) or "gen_ai." (new).
527fn is_ai_span(span_data: &Annotated<SpanData>, span_op: Option<&OperationType>) -> bool {
528    let has_ai_op = span_data
529        .value()
530        .and_then(|data| data.get_value(GEN_AI__OPERATION__NAME))
531        .is_some();
532
533    let is_ai_span_op =
534        span_op.is_some_and(|op| op.starts_with("ai.") || op.starts_with("gen_ai."));
535
536    has_ai_op || is_ai_span_op
537}
538
539#[cfg(test)]
540mod tests {
541    use std::collections::HashMap;
542
543    use relay_protocol::{FromValue, assert_annotated_snapshot};
544    use serde_json::json;
545
546    use super::*;
547    use crate::ModelMetadataEntry;
548
549    fn ai_span_with_data(data: serde_json::Value) -> Span {
550        Span {
551            op: "gen_ai.test".to_owned().into(),
552            data: SpanData::from_value(data.into()),
553            ..Default::default()
554        }
555    }
556
557    #[test]
558    fn test_has_valid_total_cost() {
559        let missing = ai_span_with_data(json!({}));
560        let invalid = ai_span_with_data(json!({"gen_ai.cost.total_tokens": false}));
561        let valid = ai_span_with_data(json!({"gen_ai.cost.total_tokens": 1.0}));
562
563        assert!(!has_valid_total_cost(missing.data.value().unwrap()));
564        assert!(!has_valid_total_cost(invalid.data.value().unwrap()));
565        assert!(has_valid_total_cost(valid.data.value().unwrap()));
566    }
567
568    #[test]
569    fn test_calculate_cost_no_tokens() {
570        let cost = calculate_costs(
571            &ModelCostV2 {
572                input_per_token: 1.0,
573                output_per_token: 1.0,
574                output_reasoning_per_token: 1.0,
575                input_cached_per_token: 1.0,
576                input_cache_write_per_token: 1.0,
577            },
578            UsedTokens::from_span_data(&SpanData::default()),
579            "test",
580            "test",
581        );
582        assert!(cost.is_none());
583    }
584
585    #[test]
586    fn test_calculate_cost_full() {
587        let cost = calculate_costs(
588            &ModelCostV2 {
589                input_per_token: 1.0,
590                output_per_token: 2.0,
591                output_reasoning_per_token: 3.0,
592                input_cached_per_token: 0.5,
593                input_cache_write_per_token: 0.75,
594            },
595            UsedTokens {
596                input_tokens: 8.0,
597                input_cached_tokens: 5.0,
598                input_cache_write_tokens: 0.0,
599                output_tokens: 15.0,
600                output_reasoning_tokens: 9.0,
601            },
602            "test",
603            "test",
604        )
605        .unwrap();
606
607        insta::assert_debug_snapshot!(cost, @r"
608        CalculatedCost {
609            input: 5.5,
610            output: 39.0,
611            cache_read_input: 2.5,
612            cache_creation_input: 0.0,
613            reasoning_output: 27.0,
614        }
615        ");
616    }
617
618    #[test]
619    fn test_calculate_cost_no_reasoning_cost() {
620        let cost = calculate_costs(
621            &ModelCostV2 {
622                input_per_token: 1.0,
623                output_per_token: 2.0,
624                // Should fallback to output token cost for reasoning.
625                output_reasoning_per_token: 0.0,
626                input_cached_per_token: 0.5,
627                input_cache_write_per_token: 0.0,
628            },
629            UsedTokens {
630                input_tokens: 8.0,
631                input_cached_tokens: 5.0,
632                input_cache_write_tokens: 0.0,
633                output_tokens: 15.0,
634                output_reasoning_tokens: 9.0,
635            },
636            "test",
637            "test",
638        )
639        .unwrap();
640
641        insta::assert_debug_snapshot!(cost, @r"
642        CalculatedCost {
643            input: 5.5,
644            output: 30.0,
645            cache_read_input: 2.5,
646            cache_creation_input: 0.0,
647            reasoning_output: 18.0,
648        }
649        ");
650    }
651
652    /// This test shows it is possible to produce negative costs if tokens are not aligned properly.
653    ///
654    /// The behaviour was desired when initially implemented.
655    #[test]
656    fn test_calculate_cost_negative() {
657        let cost = calculate_costs(
658            &ModelCostV2 {
659                input_per_token: 2.0,
660                output_per_token: 2.0,
661                output_reasoning_per_token: 1.0,
662                input_cached_per_token: 1.0,
663                input_cache_write_per_token: 1.5,
664            },
665            UsedTokens {
666                input_tokens: 1.0,
667                input_cached_tokens: 11.0,
668                input_cache_write_tokens: 0.0,
669                output_tokens: 1.0,
670                output_reasoning_tokens: 9.0,
671            },
672            "test",
673            "test",
674        )
675        .unwrap();
676
677        insta::assert_debug_snapshot!(cost, @r"
678        CalculatedCost {
679            input: -9.0,
680            output: -7.0,
681            cache_read_input: 11.0,
682            cache_creation_input: 0.0,
683            reasoning_output: 9.0,
684        }
685        ");
686    }
687
688    #[test]
689    fn test_calculate_cost_with_cache_writes() {
690        let cost = calculate_costs(
691            &ModelCostV2 {
692                input_per_token: 1.0,
693                output_per_token: 2.0,
694                output_reasoning_per_token: 3.0,
695                input_cached_per_token: 0.5,
696                input_cache_write_per_token: 0.75,
697            },
698            UsedTokens {
699                input_tokens: 100.0,
700                input_cached_tokens: 20.0,
701                input_cache_write_tokens: 30.0,
702                output_tokens: 50.0,
703                output_reasoning_tokens: 10.0,
704            },
705            "test",
706            "test",
707        )
708        .unwrap();
709
710        // input: (100 - 20 - 30) * 1.0 + 20 * 0.5 + 30 * 0.75 = 50 + 10 + 22.5 = 82.5
711        //   (cache-write tokens are billed once at the cache-write rate, not also at
712        //    the standard input rate). output: 40 * 2.0 + 10 * 3.0 = 110.0
713        insta::assert_debug_snapshot!(cost, @r"
714        CalculatedCost {
715            input: 82.5,
716            output: 110.0,
717            cache_read_input: 10.0,
718            cache_creation_input: 22.5,
719            reasoning_output: 30.0,
720        }
721        ");
722    }
723
724    #[test]
725    fn test_existing_cost_is_not_overwritten() {
726        let mut span = ai_span_with_data(json!({
727            "gen_ai.response.model": "claude-2.1",
728            "gen_ai.usage.input_tokens": 1000.0,
729            "gen_ai.cost.input_tokens": 99.0,
730            "gen_ai.cost.total_tokens": 123.0,
731        }));
732
733        enrich_ai_span(&mut span, Some(&metadata_with_context_size()));
734
735        let data = span.data.value().unwrap();
736        assert_eq!(
737            data.get_value(GEN_AI__COST__TOTAL_TOKENS)
738                .and_then(Value::as_f64),
739            Some(123.0)
740        );
741        assert!(data.get_value(GEN_AI__COST__OUTPUT_TOKENS).is_none());
742    }
743
744    #[test]
745    fn test_calculate_cost_backward_compatibility_no_cache_write() {
746        // Test that cost calculation works when cache_write field is missing (backward compatibility)
747        let span_data = SpanData::from([
748            (
749                GEN_AI__USAGE__INPUT_TOKENS.to_owned(),
750                Annotated::new(100.0.into()),
751            ),
752            (
753                GEN_AI__USAGE__CACHE_READ__INPUT_TOKENS.to_owned(),
754                Annotated::new(20.0.into()),
755            ),
756            (
757                GEN_AI__USAGE__OUTPUT_TOKENS.to_owned(),
758                Annotated::new(50.0.into()),
759            ),
760        ]);
761
762        let tokens = UsedTokens::from_span_data(&span_data);
763
764        // Verify cache_write_tokens defaults to 0.0
765        assert_eq!(tokens.input_cache_write_tokens, 0.0);
766
767        let cost = calculate_costs(
768            &ModelCostV2 {
769                input_per_token: 1.0,
770                output_per_token: 2.0,
771                output_reasoning_per_token: 0.0,
772                input_cached_per_token: 0.5,
773                input_cache_write_per_token: 0.75,
774            },
775            tokens,
776            "test",
777            "test",
778        )
779        .unwrap();
780
781        // Cost should be calculated without cache_write_tokens
782        // input: (100 - 20) * 1.0 + 20 * 0.5 + 0 * 0.75 = 80 + 10 + 0 = 90
783        // output: 50 * 2.0 = 100
784        insta::assert_debug_snapshot!(cost, @r"
785        CalculatedCost {
786            input: 90.0,
787            output: 100.0,
788            cache_read_input: 10.0,
789            cache_creation_input: 0.0,
790            reasoning_output: 0.0,
791        }
792        ");
793    }
794
795    /// Test that the AI operation type is inferred from a gen_ai.operation.name attribute.
796    #[test]
797    fn test_infer_ai_operation_type_from_gen_ai_operation_name() {
798        let mut span = ai_span_with_data(json!({
799            "gen_ai.operation.name": "invoke_agent"
800        }));
801
802        enrich_ai_span(&mut span, None);
803
804        assert_annotated_snapshot!(&span.data, @r#"
805        {
806          "gen_ai.operation.name": "invoke_agent",
807          "gen_ai.operation.type": "agent"
808        }
809        "#);
810    }
811
812    /// Test that the AI operation type is inferred from a span.op attribute.
813    #[test]
814    fn test_infer_ai_operation_type_from_span_op() {
815        let mut span = Span {
816            op: "gen_ai.invoke_agent".to_owned().into(),
817            ..Default::default()
818        };
819
820        enrich_ai_span(&mut span, None);
821
822        assert_annotated_snapshot!(span.data, @r#"
823        {
824          "gen_ai.operation.type": "agent"
825        }
826        "#);
827    }
828
829    /// Test that the AI operation type is inferred from a fallback.
830    #[test]
831    fn test_infer_ai_operation_type_from_fallback() {
832        let mut span = ai_span_with_data(json!({
833            "gen_ai.operation.name": "embeddings"
834        }));
835
836        enrich_ai_span(&mut span, None);
837
838        assert_annotated_snapshot!(&span.data, @r#"
839        {
840          "gen_ai.operation.name": "embeddings",
841          "gen_ai.operation.type": "ai_client"
842        }
843        "#);
844    }
845
846    /// Test that the response model is defaulted to the request model if not set.
847    #[test]
848    fn test_default_response_model_from_request_model() {
849        let mut span = ai_span_with_data(json!({
850            "gen_ai.request.model": "gpt-4",
851        }));
852
853        enrich_ai_span(&mut span, None);
854
855        assert_annotated_snapshot!(&span.data, @r#"
856        {
857          "gen_ai.operation.type": "ai_client",
858          "gen_ai.request.model": "gpt-4",
859          "gen_ai.response.model": "gpt-4"
860        }
861        "#);
862    }
863
864    /// Test that the response model is defaulted to the request model if not set.
865    #[test]
866    fn test_default_response_model_not_overridden() {
867        let mut span = ai_span_with_data(json!({
868            "gen_ai.request.model": "gpt-4",
869            "gen_ai.response.model": "gpt-4-abcd",
870        }));
871
872        enrich_ai_span(&mut span, None);
873
874        assert_annotated_snapshot!(&span.data, @r#"
875        {
876          "gen_ai.operation.type": "ai_client",
877          "gen_ai.request.model": "gpt-4",
878          "gen_ai.response.model": "gpt-4-abcd"
879        }
880        "#);
881    }
882
883    /// Test that gen_ai.agent.name is defaulted from gen_ai.function_id.
884    #[test]
885    fn test_default_agent_name_from_function_id() {
886        let mut span = ai_span_with_data(json!({
887            "gen_ai.function_id": "my-agent",
888        }));
889
890        enrich_ai_span(&mut span, None);
891
892        assert_annotated_snapshot!(&span.data, @r#"
893        {
894          "gen_ai.agent.name": "my-agent",
895          "gen_ai.function_id": "my-agent",
896          "gen_ai.operation.type": "ai_client"
897        }
898        "#);
899    }
900
901    /// Test that gen_ai.agent.name is not overridden when already set.
902    #[test]
903    fn test_default_agent_name_not_overridden() {
904        let mut span = ai_span_with_data(json!({
905            "gen_ai.function_id": "my-function",
906            "gen_ai.agent.name": "my-agent",
907        }));
908
909        enrich_ai_span(&mut span, None);
910
911        assert_annotated_snapshot!(&span.data, @r#"
912        {
913          "gen_ai.agent.name": "my-agent",
914          "gen_ai.function_id": "my-function",
915          "gen_ai.operation.type": "ai_client"
916        }
917        "#);
918    }
919
920    /// Test that an AI span is detected from a gen_ai.operation.name attribute.
921    #[test]
922    fn test_is_ai_span_from_gen_ai_operation_name() {
923        let mut span_data = Annotated::default();
924        span_data
925            .get_or_insert_with(SpanData::default)
926            .other
927            .insert(
928                GEN_AI__OPERATION__NAME.to_owned(),
929                Annotated::new(Value::String("chat".into())),
930            );
931        assert!(is_ai_span(&span_data, None));
932    }
933
934    /// Test that an AI span is detected from a span.op starting with "ai.".
935    #[test]
936    fn test_is_ai_span_from_span_op_ai() {
937        let span_op: OperationType = "ai.chat".into();
938        assert!(is_ai_span(&Annotated::default(), Some(&span_op)));
939    }
940
941    /// Test that an AI span is detected from a span.op starting with "gen_ai.".
942    #[test]
943    fn test_is_ai_span_from_span_op_gen_ai() {
944        let span_op: OperationType = "gen_ai.chat".into();
945        assert!(is_ai_span(&Annotated::default(), Some(&span_op)));
946    }
947
948    /// Test that a non-AI span is detected.
949    #[test]
950    fn test_is_ai_span_negative() {
951        assert!(!is_ai_span(&Annotated::default(), None));
952    }
953
954    /// Test enrich_ai_event_data with invoke_agent in trace context and a chat child span.
955    #[test]
956    fn test_enrich_ai_event_data_invoke_agent_trace_with_chat_span() {
957        let event_json = r#"{
958            "type": "transaction",
959            "timestamp": 1234567892.0,
960            "start_timestamp": 1234567889.0,
961            "contexts": {
962                "trace": {
963                    "op": "gen_ai.invoke_agent",
964                    "trace_id": "12345678901234567890123456789012",
965                    "span_id": "1234567890123456",
966                    "data": {
967                        "gen_ai.operation.name": "gen_ai.invoke_agent",
968                        "gen_ai.usage.input_tokens": 500,
969                        "gen_ai.usage.output_tokens": 200
970                    }
971                }
972            },
973            "spans": [
974                {
975                    "op": "gen_ai.chat.completions",
976                    "span_id": "1234567890123457",
977                    "start_timestamp": 1234567889.5,
978                    "timestamp": 1234567890.5,
979                    "data": {
980                        "gen_ai.operation.name": "chat",
981                        "gen_ai.usage.input_tokens": 100,
982                        "gen_ai.usage.output_tokens": 50
983                    }
984                }
985            ]
986        }"#;
987
988        let mut annotated_event: Annotated<Event> = Annotated::from_json(event_json).unwrap();
989        let event = annotated_event.value_mut().as_mut().unwrap();
990
991        enrich_ai_event_data(event, None);
992
993        assert_annotated_snapshot!(&annotated_event, @r#"
994        {
995          "type": "transaction",
996          "timestamp": 1234567892.0,
997          "start_timestamp": 1234567889.0,
998          "contexts": {
999            "trace": {
1000              "trace_id": "12345678901234567890123456789012",
1001              "span_id": "1234567890123456",
1002              "op": "gen_ai.invoke_agent",
1003              "data": {
1004                "gen_ai.operation.name": "gen_ai.invoke_agent",
1005                "gen_ai.operation.type": "agent",
1006                "gen_ai.usage.input_tokens": 500,
1007                "gen_ai.usage.output_tokens": 200,
1008                "gen_ai.usage.total_tokens": 700.0
1009              },
1010              "type": "trace"
1011            }
1012          },
1013          "spans": [
1014            {
1015              "timestamp": 1234567890.5,
1016              "start_timestamp": 1234567889.5,
1017              "op": "gen_ai.chat.completions",
1018              "span_id": "1234567890123457",
1019              "data": {
1020                "gen_ai.operation.name": "chat",
1021                "gen_ai.operation.type": "ai_client",
1022                "gen_ai.usage.input_tokens": 100,
1023                "gen_ai.usage.output_tokens": 50,
1024                "gen_ai.usage.total_tokens": 150.0
1025              }
1026            }
1027          ]
1028        }
1029        "#);
1030    }
1031
1032    /// Test enrich_ai_event_data with non-AI trace context, invoke_agent parent span, and chat child span.
1033    #[test]
1034    fn test_enrich_ai_event_data_nested_agent_and_chat_spans() {
1035        let event_json = r#"{
1036            "type": "transaction",
1037            "timestamp": 1234567892.0,
1038            "start_timestamp": 1234567889.0,
1039            "contexts": {
1040                "trace": {
1041                    "op": "http.server",
1042                    "trace_id": "12345678901234567890123456789012",
1043                    "span_id": "1234567890123456"
1044                }
1045            },
1046            "spans": [
1047                {
1048                    "op": "gen_ai.invoke_agent",
1049                    "span_id": "1234567890123457",
1050                    "parent_span_id": "1234567890123456",
1051                    "start_timestamp": 1234567889.5,
1052                    "timestamp": 1234567891.5,
1053                    "data": {
1054                        "gen_ai.operation.name": "invoke_agent",
1055                        "gen_ai.usage.input_tokens": 500,
1056                        "gen_ai.usage.output_tokens": 200
1057                    }
1058                },
1059                {
1060                    "op": "gen_ai.chat.completions",
1061                    "span_id": "1234567890123458",
1062                    "parent_span_id": "1234567890123457",
1063                    "start_timestamp": 1234567890.0,
1064                    "timestamp": 1234567891.0,
1065                    "data": {
1066                        "gen_ai.operation.name": "chat",
1067                        "gen_ai.usage.input_tokens": 100,
1068                        "gen_ai.usage.output_tokens": 50
1069                    }
1070                }
1071            ]
1072        }"#;
1073
1074        let mut annotated_event: Annotated<Event> = Annotated::from_json(event_json).unwrap();
1075        let event = annotated_event.value_mut().as_mut().unwrap();
1076
1077        enrich_ai_event_data(event, None);
1078
1079        assert_annotated_snapshot!(&annotated_event, @r#"
1080        {
1081          "type": "transaction",
1082          "timestamp": 1234567892.0,
1083          "start_timestamp": 1234567889.0,
1084          "contexts": {
1085            "trace": {
1086              "trace_id": "12345678901234567890123456789012",
1087              "span_id": "1234567890123456",
1088              "op": "http.server",
1089              "type": "trace"
1090            }
1091          },
1092          "spans": [
1093            {
1094              "timestamp": 1234567891.5,
1095              "start_timestamp": 1234567889.5,
1096              "op": "gen_ai.invoke_agent",
1097              "span_id": "1234567890123457",
1098              "parent_span_id": "1234567890123456",
1099              "data": {
1100                "gen_ai.operation.name": "invoke_agent",
1101                "gen_ai.operation.type": "agent",
1102                "gen_ai.usage.input_tokens": 500,
1103                "gen_ai.usage.output_tokens": 200,
1104                "gen_ai.usage.total_tokens": 700.0
1105              }
1106            },
1107            {
1108              "timestamp": 1234567891.0,
1109              "start_timestamp": 1234567890.0,
1110              "op": "gen_ai.chat.completions",
1111              "span_id": "1234567890123458",
1112              "parent_span_id": "1234567890123457",
1113              "data": {
1114                "gen_ai.operation.name": "chat",
1115                "gen_ai.operation.type": "ai_client",
1116                "gen_ai.usage.input_tokens": 100,
1117                "gen_ai.usage.output_tokens": 50,
1118                "gen_ai.usage.total_tokens": 150.0
1119              }
1120            }
1121          ]
1122        }
1123        "#);
1124    }
1125
1126    /// Test enrich_ai_event_data with legacy measurements and span op for operation type.
1127    #[test]
1128    fn test_enrich_ai_event_data_legacy_measurements_and_span_op() {
1129        let event_json = r#"{
1130            "type": "transaction",
1131            "timestamp": 1234567892.0,
1132            "start_timestamp": 1234567889.0,
1133            "contexts": {
1134                "trace": {
1135                    "op": "http.server",
1136                    "trace_id": "12345678901234567890123456789012",
1137                    "span_id": "1234567890123456"
1138                }
1139            },
1140            "spans": [
1141                {
1142                    "op": "gen_ai.invoke_agent",
1143                    "span_id": "1234567890123457",
1144                    "parent_span_id": "1234567890123456",
1145                    "start_timestamp": 1234567889.5,
1146                    "timestamp": 1234567891.5,
1147                    "measurements": {
1148                        "ai_prompt_tokens_used": {"value": 500.0},
1149                        "ai_completion_tokens_used": {"value": 200.0}
1150                    }
1151                },
1152                {
1153                    "op": "ai.chat_completions.create.langchain.ChatOpenAI",
1154                    "span_id": "1234567890123458",
1155                    "parent_span_id": "1234567890123457",
1156                    "start_timestamp": 1234567890.0,
1157                    "timestamp": 1234567891.0,
1158                    "measurements": {
1159                        "ai_prompt_tokens_used": {"value": 100.0},
1160                        "ai_completion_tokens_used": {"value": 50.0}
1161                    }
1162                }
1163            ]
1164        }"#;
1165
1166        let mut annotated_event: Annotated<Event> = Annotated::from_json(event_json).unwrap();
1167        let event = annotated_event.value_mut().as_mut().unwrap();
1168
1169        enrich_ai_event_data(event, None);
1170
1171        assert_annotated_snapshot!(&annotated_event, @r#"
1172        {
1173          "type": "transaction",
1174          "timestamp": 1234567892.0,
1175          "start_timestamp": 1234567889.0,
1176          "contexts": {
1177            "trace": {
1178              "trace_id": "12345678901234567890123456789012",
1179              "span_id": "1234567890123456",
1180              "op": "http.server",
1181              "type": "trace"
1182            }
1183          },
1184          "spans": [
1185            {
1186              "timestamp": 1234567891.5,
1187              "start_timestamp": 1234567889.5,
1188              "op": "gen_ai.invoke_agent",
1189              "span_id": "1234567890123457",
1190              "parent_span_id": "1234567890123456",
1191              "data": {
1192                "gen_ai.operation.type": "agent",
1193                "gen_ai.usage.input_tokens": 500.0,
1194                "gen_ai.usage.output_tokens": 200.0,
1195                "gen_ai.usage.total_tokens": 700.0
1196              },
1197              "measurements": {
1198                "ai_completion_tokens_used": {
1199                  "value": 200.0
1200                },
1201                "ai_prompt_tokens_used": {
1202                  "value": 500.0
1203                }
1204              }
1205            },
1206            {
1207              "timestamp": 1234567891.0,
1208              "start_timestamp": 1234567890.0,
1209              "op": "ai.chat_completions.create.langchain.ChatOpenAI",
1210              "span_id": "1234567890123458",
1211              "parent_span_id": "1234567890123457",
1212              "data": {
1213                "gen_ai.operation.type": "ai_client",
1214                "gen_ai.usage.input_tokens": 100.0,
1215                "gen_ai.usage.output_tokens": 50.0,
1216                "gen_ai.usage.total_tokens": 150.0
1217              },
1218              "measurements": {
1219                "ai_completion_tokens_used": {
1220                  "value": 50.0
1221                },
1222                "ai_prompt_tokens_used": {
1223                  "value": 100.0
1224                }
1225              }
1226            }
1227          ]
1228        }
1229        "#);
1230    }
1231
1232    fn metadata_with_context_size() -> ModelMetadata {
1233        ModelMetadata {
1234            version: 1,
1235            models: HashMap::from([(
1236                "claude-2.1".parse().unwrap(),
1237                ModelMetadataEntry {
1238                    costs: Some(ModelCostV2 {
1239                        input_per_token: 0.01,
1240                        output_per_token: 0.02,
1241                        output_reasoning_per_token: 0.0,
1242                        input_cached_per_token: 0.0,
1243                        input_cache_write_per_token: 0.0,
1244                    }),
1245                    context_size: Some(100_000),
1246                },
1247            )]),
1248        }
1249    }
1250
1251    #[test]
1252    fn test_context_utilization_with_total_tokens() {
1253        let mut span = Span {
1254            op: "gen_ai.test".to_owned().into(),
1255            data: SpanData::from_value(
1256                json!({
1257                    "gen_ai.response.model": "claude-2.1",
1258                    "gen_ai.usage.input_tokens": 30000.0,
1259                    "gen_ai.usage.output_tokens": 12000.0,
1260                    "gen_ai.usage.total_tokens": 42000.0,
1261                })
1262                .into(),
1263            ),
1264            ..Default::default()
1265        };
1266
1267        enrich_ai_span(&mut span, Some(&metadata_with_context_size()));
1268
1269        let data = span.data.value().unwrap();
1270        assert_eq!(
1271            data.get_value(GEN_AI__CONTEXT__WINDOW_SIZE)
1272                .and_then(Value::as_f64),
1273            Some(100_000.0)
1274        );
1275        assert_eq!(
1276            data.get_value(GEN_AI__CONTEXT__UTILIZATION)
1277                .and_then(Value::as_f64),
1278            Some(0.42)
1279        );
1280    }
1281
1282    #[test]
1283    fn test_context_utilization_no_context_size() {
1284        let metadata = ModelMetadata {
1285            version: 1,
1286            models: HashMap::from([(
1287                "claude-2.1".parse().unwrap(),
1288                ModelMetadataEntry {
1289                    costs: None,
1290                    context_size: None,
1291                },
1292            )]),
1293        };
1294
1295        let mut span = Span {
1296            op: "gen_ai.test".to_owned().into(),
1297            data: SpanData::from_value(
1298                json!({
1299                    "gen_ai.response.model": "claude-2.1",
1300                    "gen_ai.usage.total_tokens": 1000.0,
1301                })
1302                .into(),
1303            ),
1304            ..Default::default()
1305        };
1306
1307        enrich_ai_span(&mut span, Some(&metadata));
1308
1309        let data = span.data.value().unwrap();
1310        assert!(data.get_value(GEN_AI__CONTEXT__WINDOW_SIZE).is_none());
1311        assert!(data.get_value(GEN_AI__CONTEXT__UTILIZATION).is_none());
1312    }
1313
1314    #[test]
1315    fn test_context_utilization_no_total_tokens() {
1316        let mut span = Span {
1317            op: "gen_ai.test".to_owned().into(),
1318            data: SpanData::from_value(
1319                json!({
1320                    "gen_ai.response.model": "claude-2.1",
1321                })
1322                .into(),
1323            ),
1324            ..Default::default()
1325        };
1326
1327        enrich_ai_span(&mut span, Some(&metadata_with_context_size()));
1328
1329        let data = span.data.value().unwrap();
1330        // window_size should still be set even without tokens.
1331        assert_eq!(
1332            data.get_value(GEN_AI__CONTEXT__WINDOW_SIZE)
1333                .and_then(Value::as_f64),
1334            Some(100_000.0)
1335        );
1336        // But utilization cannot be computed without total_tokens.
1337        assert!(data.get_value(GEN_AI__CONTEXT__UTILIZATION).is_none());
1338    }
1339
1340    #[test]
1341    fn test_context_utilization_unknown_model() {
1342        let mut span = Span {
1343            op: "gen_ai.test".to_owned().into(),
1344            data: SpanData::from_value(
1345                json!({
1346                    "gen_ai.response.model": "unknown-model",
1347                    "gen_ai.usage.total_tokens": 1000.0,
1348                })
1349                .into(),
1350            ),
1351            ..Default::default()
1352        };
1353
1354        enrich_ai_span(&mut span, Some(&metadata_with_context_size()));
1355
1356        let data = span.data.value().unwrap();
1357        assert!(data.get_value(GEN_AI__CONTEXT__WINDOW_SIZE).is_none());
1358        assert!(data.get_value(GEN_AI__CONTEXT__UTILIZATION).is_none());
1359    }
1360}