diff --git a/README.md b/README.md index bf63e628..e5544561 100644 --- a/README.md +++ b/README.md @@ -26,6 +26,7 @@ - [*Romeo and Juliet* Full Text Extraction](#romeo-and-juliet-full-text-extraction) - [Medication Extraction](#medication-extraction) - [Radiology Report Structuring: RadExtract](#radiology-report-structuring-radextract) + - [Tracking Token Usage & API Calls](#tracking-token-usage--api-calls) - [Community Providers](#community-providers) - [Contributing](#contributing) - [Testing](#testing) @@ -175,6 +176,52 @@ result = lx.extract( This approach can extract hundreds of entities from full novels while maintaining high accuracy. The interactive visualization seamlessly handles large result sets, making it easy to explore hundreds of entities from the output JSONL file. **[See the full *Romeo and Juliet* extraction example →](https://github.com/google/langextract/blob/main/docs/examples/longer_text_example.md)** for detailed results and performance insights. +### Tracking Token Usage & API Calls + +By default, the returned `AnnotatedDocument` contains aggregated token usage and API call metrics inside its `metadata` dictionary: + +```python +result = lx.extract( + text_or_documents=input_text, + prompt_description=prompt, + examples=examples, + model_id="gemini-3.5-flash", +) + +print(result.metadata) +# Output: +# { +# 'token_usage': {'prompt_tokens': 197, 'completion_tokens': 82, 'total_tokens': 279}, +# 'api_calls': 1 +# } +``` + +> **Note:** The total number of `api_calls` is equal to `(number of chunks) * (extraction_passes)`. +> +> **Warning:** Enabling `track_api_call_details=True` on extremely large runs (with thousands of chunks or passes) can consume significant memory because details for every single API call are stored in memory. Only use it when detailed observability is required. + +For detailed observability (e.g., tracking the exact token cost of each chunk or extraction pass), you can opt-in to detailed call-level tracking by setting `track_api_call_details=True`: + +```python +result = lx.extract( + text_or_documents=input_text, + prompt_description=prompt, + examples=examples, + model_id="gemini-3.5-flash", + track_api_call_details=True, +) + +print(result.metadata["api_call_details"]) +# Output: +# [ +# { +# 'pass_index': 0, +# 'chunk_index': 0, +# 'token_usage': {'prompt_tokens': 197, 'completion_tokens': 82, 'total_tokens': 279} +# } +# ] +``` + ### Vertex AI Batch Processing Save costs on large-scale tasks by enabling Vertex AI Batch API with diff --git a/langextract/annotation.py b/langextract/annotation.py index d77ab178..abc08dbd 100644 --- a/langextract/annotation.py +++ b/langextract/annotation.py @@ -217,6 +217,7 @@ def annotate_documents( context_window_chars: int | None = None, show_progress: bool = True, tokenizer: tokenizer_lib.Tokenizer | None = None, + track_api_call_details: bool = False, **kwargs, ) -> Iterator[data.AnnotatedDocument]: """Annotates a sequence of documents with NLP extractions. @@ -244,6 +245,9 @@ def annotate_documents( resolution across chunk boundaries. Defaults to None (disabled). show_progress: Whether to show progress bar. Defaults to True. tokenizer: Optional tokenizer to use. If None, uses default tokenizer. + track_api_call_details: Whether to track detailed tokens of individual API calls. + Warning: Enabling this on extremely large runs with thousands of chunks/passes can consume + significant memory. **kwargs: Additional arguments passed to LanguageModel.infer and Resolver. @@ -266,6 +270,7 @@ def annotate_documents( show_progress, context_window_chars=context_window_chars, tokenizer=tokenizer, + track_api_call_details=track_api_call_details, **kwargs, ) else: @@ -279,6 +284,7 @@ def annotate_documents( show_progress, context_window_chars=context_window_chars, tokenizer=tokenizer, + track_api_call_details=track_api_call_details, **kwargs, ) @@ -293,6 +299,8 @@ def _annotate_documents_single_pass( context_window_chars: int | None = None, tokenizer: tokenizer_lib.Tokenizer | None = None, suppress_parse_errors: bool = False, + track_api_call_details: bool = False, + pass_num: int = 0, **kwargs, ) -> Iterator[data.AnnotatedDocument]: """Single-pass annotation with stable ordering and streaming emission. @@ -309,6 +317,14 @@ def _annotate_documents_single_pass( per_doc: DefaultDict[str, list[data.Extraction]] = collections.defaultdict( list ) + doc_usage = collections.defaultdict( + lambda: {"prompt_tokens": 0, "completion_tokens": 0, "total_tokens": 0} + ) + # doc_successful_api_calls tracks the number of successful inference API requests + # that completed and returned a scored output. + doc_successful_api_calls = collections.defaultdict(int) + doc_api_call_details = collections.defaultdict(list) + chunk_counters = collections.defaultdict(int) next_emit_idx = 0 def _capture_docs(src: Iterable[data.Document]) -> Iterator[data.Document]: @@ -336,13 +352,26 @@ def _emit_docs_iter( limit = max(0, len(doc_order) - 1) if keep_last_doc else len(doc_order) while next_emit_idx < limit: document_id = doc_order[next_emit_idx] + metadata = { + "token_usage": doc_usage.get(document_id), + "api_calls": doc_successful_api_calls.get(document_id, 0), + } + if track_api_call_details: + metadata["api_call_details"] = doc_api_call_details.get( + document_id, [] + ) yield data.AnnotatedDocument( document_id=document_id, extractions=per_doc.get(document_id, []), text=doc_text_by_id.get(document_id, ""), + metadata=metadata, ) per_doc.pop(document_id, None) doc_text_by_id.pop(document_id, None) + doc_usage.pop(document_id, None) + doc_successful_api_calls.pop(document_id, None) + doc_api_call_details.pop(document_id, None) + chunk_counters.pop(document_id, None) next_emit_idx += 1 chunk_iter = _document_chunk_iterator( @@ -401,8 +430,37 @@ def _emit_docs_iter( "No scored outputs from language model." ) + doc_id = text_chunk.document_id + doc_successful_api_calls[doc_id] += 1 + scored_output = scored_outputs[0] + usage = scored_output.token_usage + chunk_idx = chunk_counters[doc_id] + chunk_counters[doc_id] += 1 + + if usage is not None: + for key in ("prompt_tokens", "completion_tokens", "total_tokens"): + val = getattr(usage, key, None) + if isinstance(val, (int, float)): + doc_usage[doc_id][key] += val + + if track_api_call_details: + detail = { + "pass_index": pass_num, + "chunk_index": chunk_idx, + "token_usage": { + "prompt_tokens": usage.prompt_tokens if usage else None, + "completion_tokens": ( + usage.completion_tokens if usage else None + ), + "total_tokens": usage.total_tokens if usage else None, + }, + } + if scored_output.request_id is not None: + detail["request_id"] = scored_output.request_id + doc_api_call_details[doc_id].append(detail) + resolved_extractions = resolver.resolve( - scored_outputs[0].output, + scored_output.output, debug=debug, suppress_parse_errors=suppress_parse_errors, **kwargs, @@ -455,6 +513,7 @@ def _annotate_documents_sequential_passes( show_progress: bool = True, context_window_chars: int | None = None, tokenizer: tokenizer_lib.Tokenizer | None = None, + track_api_call_details: bool = False, **kwargs, ) -> Iterator[data.AnnotatedDocument]: """Sequential extraction passes logic for improved recall.""" @@ -469,6 +528,10 @@ def _annotate_documents_sequential_passes( document_extractions_by_pass: dict[str, list[list[data.Extraction]]] = {} document_texts: dict[str, str] = {} + document_usage = {} + document_api_calls = {} + document_api_call_details = {} + # Preserve text up-front so we can emit documents even if later passes # produce no extractions. for _doc in document_list: @@ -488,18 +551,43 @@ def _annotate_documents_sequential_passes( show_progress=show_progress if pass_num == 0 else False, context_window_chars=context_window_chars, tokenizer=tokenizer, + track_api_call_details=track_api_call_details, + pass_num=pass_num, **kwargs, ): doc_id = annotated_doc.document_id if doc_id not in document_extractions_by_pass: document_extractions_by_pass[doc_id] = [] - # Keep first-seen text (already pre-filled above). + document_usage[doc_id] = { + "prompt_tokens": 0, + "completion_tokens": 0, + "total_tokens": 0, + } + document_api_calls[doc_id] = 0 + document_api_call_details[doc_id] = [] document_extractions_by_pass[doc_id].append( annotated_doc.extractions or [] ) + if annotated_doc.metadata: + meta = annotated_doc.metadata + usage = meta.get("token_usage") + if usage: + document_usage[doc_id]["prompt_tokens"] += ( + usage.get("prompt_tokens") or 0 + ) + document_usage[doc_id]["completion_tokens"] += ( + usage.get("completion_tokens") or 0 + ) + document_usage[doc_id]["total_tokens"] += ( + usage.get("total_tokens") or 0 + ) + document_api_calls[doc_id] += meta.get("api_calls", 0) + if "api_call_details" in meta: + document_api_call_details[doc_id].extend(meta["api_call_details"]) + # Emit results strictly in original input order. for doc in document_list: doc_id = doc.document_id @@ -521,10 +609,18 @@ def _annotate_documents_sequential_passes( len(merged_extractions), ) + metadata = { + "token_usage": document_usage.get(doc_id), + "api_calls": document_api_calls.get(doc_id, 0), + } + if track_api_call_details: + metadata["api_call_details"] = document_api_call_details.get(doc_id, []) + yield data.AnnotatedDocument( document_id=doc_id, extractions=merged_extractions, text=document_texts.get(doc_id, doc.text or ""), + metadata=metadata, ) logging.info("Sequential extraction passes completed.") @@ -541,6 +637,7 @@ def annotate_text( context_window_chars: int | None = None, show_progress: bool = True, tokenizer: tokenizer_lib.Tokenizer | None = None, + track_api_call_details: bool = False, **kwargs, ) -> data.AnnotatedDocument: """Annotates text with NLP extractions for text input. @@ -562,6 +659,9 @@ def annotate_text( (disabled). show_progress: Whether to show progress bar. Defaults to True. tokenizer: Optional tokenizer instance. + track_api_call_details: Whether to track detailed tokens of individual API calls. + Warning: Enabling this on extremely large runs with thousands of chunks/passes can consume + significant memory. **kwargs: Additional arguments for inference and resolver_lib. Returns: @@ -593,6 +693,7 @@ def annotate_text( context_window_chars=context_window_chars, show_progress=show_progress, tokenizer=tokenizer, + track_api_call_details=track_api_call_details, **kwargs, ) ) @@ -622,4 +723,5 @@ def annotate_text( document_id=annotations[0].document_id, extractions=annotations[0].extractions, text=annotations[0].text, + metadata=annotations[0].metadata, ) diff --git a/langextract/core/data.py b/langextract/core/data.py index 3d680c8e..6456bcc2 100644 --- a/langextract/core/data.py +++ b/langextract/core/data.py @@ -17,6 +17,7 @@ import dataclasses import enum +from typing import Any import uuid from langextract.core import tokenizer @@ -212,10 +213,14 @@ class AnnotatedDocument: extractions: List of extractions in the document. text: Raw text representation of the document. tokenized_text: Tokenized text of the document, computed from `text`. + metadata: Metadata dict (e.g. token_usage, api_calls) for the document. """ extractions: list[Extraction] | None = None text: str | None = None + metadata: dict[str, Any] | None = dataclasses.field( + default=None, compare=False + ) _document_id: str | None = dataclasses.field( default=None, init=False, repr=False, compare=False ) @@ -229,9 +234,11 @@ def __init__( document_id: str | None = None, extractions: list[Extraction] | None = None, text: str | None = None, + metadata: dict[str, Any] | None = None, ): self.extractions = extractions self.text = text + self.metadata = metadata self._document_id = document_id @property diff --git a/langextract/core/types.py b/langextract/core/types.py index ea68df23..d17e173d 100644 --- a/langextract/core/types.py +++ b/langextract/core/types.py @@ -66,12 +66,25 @@ class Constraint: constraint_type: ConstraintType = ConstraintType.NONE +@dataclasses.dataclass(frozen=True) +class TokenUsage: + """Token usage details for a model request/run.""" + + prompt_tokens: int | None = None + completion_tokens: int | None = None + total_tokens: int | None = None + + @dataclasses.dataclass(frozen=True) class ScoredOutput: """Scored output from language model inference.""" score: float | None = None output: str | None = None + token_usage: TokenUsage | None = dataclasses.field( + default=None, compare=False + ) + request_id: str | None = dataclasses.field(default=None, compare=False) def __str__(self) -> str: score_str = '-' if self.score is None else f'{self.score:.2f}' diff --git a/langextract/data_lib.py b/langextract/data_lib.py index 6b50ca8c..deeb2fce 100644 --- a/langextract/data_lib.py +++ b/langextract/data_lib.py @@ -79,6 +79,9 @@ def annotated_document_to_dict( result["document_id"] = adoc.document_id + if adoc.metadata is None: + result.pop("metadata", None) + return result @@ -121,4 +124,5 @@ def dict_to_annotated_document( extractions=[ data.Extraction(**ent) for ent in adoc_dic.get("extractions", []) ], + metadata=adoc_dic.get("metadata"), ) diff --git a/langextract/extraction.py b/langextract/extraction.py index aa60aba4..2fca12e4 100644 --- a/langextract/extraction.py +++ b/langextract/extraction.py @@ -71,6 +71,7 @@ def extract( prompt_validation_level: pv.PromptValidationLevel = pv.PromptValidationLevel.WARNING, prompt_validation_strict: bool = False, show_progress: bool = True, + track_api_call_details: bool = False, tokenizer: tokenizer_lib.Tokenizer | None = None, ) -> list[data.AnnotatedDocument] | data.AnnotatedDocument: """Extracts structured information from text. @@ -183,6 +184,11 @@ def extract( prompt_validation_strict: When True and prompt_validation_level is ERROR, raises on non-exact matches (MATCH_FUZZY, MATCH_LESSER). Defaults to False. show_progress: Whether to show progress bar during extraction. Defaults to True. + track_api_call_details: Whether to track detailed tokens of individual API calls. + Note that the total number of API calls is equal to (number of chunks) * (extraction_passes). + Warning: Enabling this on extremely large runs with thousands of chunks/passes can consume + significant memory. + Defaults to False. Returns: An AnnotatedDocument with the extracted information when input is a @@ -398,6 +404,7 @@ def extract( show_progress=show_progress, max_workers=max_workers, tokenizer=tokenizer, + track_api_call_details=track_api_call_details, **alignment_kwargs, ) return result @@ -422,6 +429,7 @@ def extract( show_progress=show_progress, max_workers=max_workers, tokenizer=tokenizer, + track_api_call_details=track_api_call_details, **alignment_kwargs, ) return list(result) diff --git a/langextract/providers/gemini.py b/langextract/providers/gemini.py index 18416820..4345a085 100644 --- a/langextract/providers/gemini.py +++ b/langextract/providers/gemini.py @@ -368,7 +368,17 @@ def _process_single_prompt( response = self._client.models.generate_content( model=self.model_id, contents=prompt, config=call_config ) - return core_types.ScoredOutput(score=1.0, output=response.text) + usage = getattr(response, 'usage_metadata', None) + token_usage = None + if usage is not None: + token_usage = core_types.TokenUsage( + prompt_tokens=getattr(usage, 'prompt_token_count', None), + completion_tokens=getattr(usage, 'candidates_token_count', None), + total_tokens=getattr(usage, 'total_token_count', None), + ) + return core_types.ScoredOutput( + score=1.0, output=response.text, token_usage=token_usage + ) except Exception as e: if attempt < self.max_retries and self._is_retryable_error(e): diff --git a/langextract/providers/ollama.py b/langextract/providers/ollama.py index 2e3470a8..c97a4532 100644 --- a/langextract/providers/ollama.py +++ b/langextract/providers/ollama.py @@ -307,7 +307,21 @@ def infer( **combined_kwargs, ) output = self._extract_response_text(response) - yield [core_types.ScoredOutput(score=1.0, output=output)] + prompt_tokens = response.get('prompt_eval_count') + completion_tokens = response.get('eval_count') + token_usage = None + if prompt_tokens is not None or completion_tokens is not None: + total_tokens = (prompt_tokens or 0) + (completion_tokens or 0) + token_usage = core_types.TokenUsage( + prompt_tokens=prompt_tokens, + completion_tokens=completion_tokens, + total_tokens=total_tokens, + ) + yield [ + core_types.ScoredOutput( + score=1.0, output=output, token_usage=token_usage + ) + ] except exceptions.InferenceError: raise except Exception as e: diff --git a/langextract/providers/openai.py b/langextract/providers/openai.py index 86974a49..1b9226cb 100644 --- a/langextract/providers/openai.py +++ b/langextract/providers/openai.py @@ -251,7 +251,22 @@ def _process_single_prompt( output_text = response.choices[0].message.content - return core_types.ScoredOutput(score=1.0, output=output_text) + usage = getattr(response, 'usage', None) + token_usage = None + if usage is not None: + token_usage = core_types.TokenUsage( + prompt_tokens=getattr(usage, 'prompt_tokens', None), + completion_tokens=getattr(usage, 'completion_tokens', None), + total_tokens=getattr(usage, 'total_tokens', None), + ) + request_id = getattr(response, 'id', None) + + return core_types.ScoredOutput( + score=1.0, + output=output_text, + token_usage=token_usage, + request_id=request_id, + ) except exceptions.InferenceConfigError: raise diff --git a/tests/token_usage_test.py b/tests/token_usage_test.py new file mode 100644 index 00000000..d6b141a3 --- /dev/null +++ b/tests/token_usage_test.py @@ -0,0 +1,284 @@ +# Copyright 2025 Google LLC. +# +# Licensed under the Apache License, Version 2.0 (the "License"); +# you may not use this file except in compliance with the License. +# You may obtain a copy of the License at +# +# http://www.apache.org/licenses/LICENSE-2.0 +# +# Unless required by applicable law or agreed to in writing, software +# distributed under the License is distributed on an "AS IS" BASIS, +# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied. +# See the License for the specific language governing permissions and +# limitations under the License. + +"""Tests for token usage and API call tracking.""" + +from unittest import mock + +from absl.testing import absltest + +from langextract import annotation +from langextract.core import data +from langextract.core import format_handler as fh +from langextract.core import types +from langextract.providers import gemini +from langextract.providers import ollama +from langextract.providers import openai + + +class TestTokenUsageTracking(absltest.TestCase): + + @mock.patch("google.genai.Client") + def test_gemini_token_usage_extraction(self, mock_client_class): + """Test that GeminiLanguageModel extracts token usage metadata.""" + mock_client = mock.Mock() + mock_client_class.return_value = mock_client + + # Simulate response carrying usage metadata + mock_response = mock.Mock() + mock_response.text = '{"extractions": []}' + mock_usage = mock.Mock() + mock_usage.prompt_token_count = 100 + mock_usage.candidates_token_count = 50 + mock_usage.total_token_count = 150 + mock_response.usage_metadata = mock_usage + + mock_client.models.generate_content.return_value = mock_response + + model = gemini.GeminiLanguageModel(api_key="test-key") + results = list(model.infer(["Test prompt"])) + + self.assertLen(results, 1) + scored_output = results[0][0] + self.assertIsNotNone(scored_output.token_usage) + self.assertEqual(scored_output.token_usage.prompt_tokens, 100) + self.assertEqual(scored_output.token_usage.completion_tokens, 50) + self.assertEqual(scored_output.token_usage.total_tokens, 150) + + @mock.patch("openai.OpenAI") + def test_openai_token_usage_extraction(self, mock_openai_class): + """Test that OpenAILanguageModel extracts token usage metadata.""" + mock_client = mock.Mock() + mock_openai_class.return_value = mock_client + + # Simulate OpenAI ChatCompletion response carrying usage and id + mock_response = mock.Mock() + mock_choice = mock.Mock() + mock_choice.message.content = '{"extractions": []}' + mock_response.choices = [mock_choice] + mock_response.id = "chatcmpl-test-id" + + mock_usage = mock.Mock() + mock_usage.prompt_tokens = 80 + mock_usage.completion_tokens = 40 + mock_usage.total_tokens = 120 + mock_response.usage = mock_usage + + mock_client.chat.completions.create.return_value = mock_response + + model = openai.OpenAILanguageModel(api_key="test-key") + results = list(model.infer(["Test prompt"])) + + self.assertLen(results, 1) + scored_output = results[0][0] + self.assertIsNotNone(scored_output.token_usage) + self.assertEqual(scored_output.token_usage.prompt_tokens, 80) + self.assertEqual(scored_output.token_usage.completion_tokens, 40) + self.assertEqual(scored_output.token_usage.total_tokens, 120) + self.assertEqual(scored_output.request_id, "chatcmpl-test-id") + + @mock.patch("requests.post") + def test_ollama_token_usage_extraction(self, mock_post): + """Test that OllamaLanguageModel extracts token usage metadata.""" + mock_response = mock.Mock() + mock_response.status_code = 200 + mock_response.json.return_value = { + "response": '{"extractions": []}', + "prompt_eval_count": 60, + "eval_count": 30, + } + mock_post.return_value = mock_response + + model = ollama.OllamaLanguageModel(model_id="gemma") + results = list(model.infer(["Test prompt"])) + + self.assertLen(results, 1) + scored_output = results[0][0] + self.assertIsNotNone(scored_output.token_usage) + self.assertEqual(scored_output.token_usage.prompt_tokens, 60) + self.assertEqual(scored_output.token_usage.completion_tokens, 30) + self.assertEqual(scored_output.token_usage.total_tokens, 90) + + def test_annotation_pipeline_aggregates_usage(self): + """Test that annotation aggregates usage and api_calls in single-pass.""" + mock_lm = mock.Mock(spec=gemini.GeminiLanguageModel) + mock_lm.requires_fence_output = False + + # We will simulate 3 chunks of inference (text of length 41 with buffer 20) + mock_lm.infer.side_effect = [ + [[ + types.ScoredOutput( + score=1.0, + output='{"extractions": [{"a": "b"}]}', + token_usage=types.TokenUsage( + prompt_tokens=10, completion_tokens=5, total_tokens=15 + ), + request_id="req-1", + ) + ]], + [[ + types.ScoredOutput( + score=1.0, + output='{"extractions": [{"c": "d"}]}', + token_usage=types.TokenUsage( + prompt_tokens=20, completion_tokens=10, total_tokens=30 + ), + request_id="req-2", + ) + ]], + [[ + types.ScoredOutput( + score=1.0, + output='{"extractions": []}', + token_usage=types.TokenUsage( + prompt_tokens=5, completion_tokens=2, total_tokens=7 + ), + request_id="req-3", + ) + ]], + ] + + mock_template = mock.Mock() + mock_template.description = "Test description" + mock_template.examples = [] + + format_handler = fh.FormatHandler() + annotator = annotation.Annotator( + language_model=mock_lm, + prompt_template=mock_template, + format_handler=format_handler, + ) + + resolver = mock.Mock() + # Return simple dummy extractions + resolver.resolve.side_effect = [ + [data.Extraction(extraction_class="value", extraction_text="b")], + [data.Extraction(extraction_class="value", extraction_text="d")], + [], + ] + resolver.align.side_effect = lambda ex, *args, **kwargs: ex + + doc = data.Document( + text="This is a test document with some chunks.", document_id="doc-1" + ) + + results = list( + annotator.annotate_documents( + documents=[doc], + resolver=resolver, + max_char_buffer=20, # Small buffer to force 3 chunks + batch_length=1, + track_api_call_details=True, + ) + ) + + self.assertLen(results, 1) + annotated_doc = results[0] + self.assertEqual(annotated_doc.document_id, "doc-1") + self.assertIsNotNone(annotated_doc.metadata) + self.assertEqual(annotated_doc.metadata["api_calls"], 3) + + usage = annotated_doc.metadata["token_usage"] + self.assertEqual(usage["prompt_tokens"], 35) + self.assertEqual(usage["completion_tokens"], 17) + self.assertEqual(usage["total_tokens"], 52) + + details = annotated_doc.metadata["api_call_details"] + self.assertLen(details, 3) + self.assertEqual(details[0]["chunk_index"], 0) + self.assertEqual(details[0]["request_id"], "req-1") + self.assertEqual(details[0]["token_usage"]["total_tokens"], 15) + self.assertEqual(details[1]["chunk_index"], 1) + self.assertEqual(details[1]["request_id"], "req-2") + self.assertEqual(details[1]["token_usage"]["total_tokens"], 30) + self.assertEqual(details[2]["chunk_index"], 2) + self.assertEqual(details[2]["request_id"], "req-3") + self.assertEqual(details[2]["token_usage"]["total_tokens"], 7) + + def test_annotation_pipeline_aggregates_usage_sequential_passes(self): + """Test that annotation aggregates usage and api_calls in sequential passes.""" + mock_lm = mock.Mock(spec=gemini.GeminiLanguageModel) + mock_lm.requires_fence_output = False + + # 2 passes, 1 chunk each + mock_lm.infer.side_effect = [ + [[ + types.ScoredOutput( + score=1.0, + output='{"extractions": []}', + token_usage=types.TokenUsage( + prompt_tokens=10, completion_tokens=5, total_tokens=15 + ), + request_id="req-pass-0", + ) + ]], + [[ + types.ScoredOutput( + score=1.0, + output='{"extractions": []}', + token_usage=types.TokenUsage( + prompt_tokens=20, completion_tokens=10, total_tokens=30 + ), + request_id="req-pass-1", + ) + ]], + ] + + mock_template = mock.Mock() + mock_template.description = "Test description" + mock_template.examples = [] + + format_handler = fh.FormatHandler() + annotator = annotation.Annotator( + language_model=mock_lm, + prompt_template=mock_template, + format_handler=format_handler, + ) + + resolver = mock.Mock() + resolver.resolve.return_value = [] + resolver.align.return_value = [] + + doc = data.Document(text="Single chunk document.", document_id="doc-1") + + results = list( + annotator.annotate_documents( + documents=[doc], + resolver=resolver, + max_char_buffer=500, + batch_length=1, + extraction_passes=2, + track_api_call_details=True, + ) + ) + + self.assertLen(results, 1) + annotated_doc = results[0] + self.assertEqual(annotated_doc.metadata["api_calls"], 2) + + usage = annotated_doc.metadata["token_usage"] + self.assertEqual(usage["prompt_tokens"], 30) + self.assertEqual(usage["completion_tokens"], 15) + self.assertEqual(usage["total_tokens"], 45) + + details = annotated_doc.metadata["api_call_details"] + self.assertLen(details, 2) + self.assertEqual(details[0]["pass_index"], 0) + self.assertEqual(details[0]["request_id"], "req-pass-0") + self.assertEqual(details[1]["pass_index"], 1) + self.assertEqual(details[1]["request_id"], "req-pass-1") + + +if __name__ == "__main__": + absltest.main()