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