Skip to content
Open
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
86 changes: 83 additions & 3 deletions src/google/adk/telemetry/sqlite_span_exporter.py
Original file line number Diff line number Diff line change
Expand Up @@ -26,10 +26,13 @@
from typing import Optional
from typing import Sequence

from opentelemetry.sdk.trace import Event
from opentelemetry.sdk.trace import ReadableSpan
from opentelemetry.sdk.trace.export import SpanExporter
from opentelemetry.sdk.trace.export import SpanExportResult
from opentelemetry.trace import SpanContext
from opentelemetry.trace import Status
from opentelemetry.trace import StatusCode
from opentelemetry.trace import TraceFlags
from opentelemetry.trace import TraceState
from opentelemetry.util.types import AttributeValue
Expand All @@ -46,10 +49,21 @@
end_time_unix_nano INTEGER,
session_id TEXT,
invocation_id TEXT,
attributes_json TEXT
attributes_json TEXT,
status_code TEXT,
status_description TEXT,
events_json TEXT
);
"""

# Columns added after the first release of this exporter. A database file
# written by an older version lacks them, so they are added on open.
_ADDED_COLUMNS = (
("status_code", "TEXT"),
("status_description", "TEXT"),
("events_json", "TEXT"),
)

_CREATE_SESSION_INDEX = """
CREATE INDEX IF NOT EXISTS spans_session_id_idx ON spans(session_id);
"""
Expand All @@ -68,8 +82,11 @@
end_time_unix_nano,
session_id,
invocation_id,
attributes_json
) VALUES (?, ?, ?, ?, ?, ?, ?, ?, ?);
attributes_json,
status_code,
status_description,
events_json
) VALUES (?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?);
"""

_DEFAULT_TIMEOUT_SECONDS = 30.0
Expand Down Expand Up @@ -102,6 +119,12 @@ def _ensure_schema(self) -> None:
with self._lock:
conn = self._get_connection()
conn.execute(_CREATE_SPANS_TABLE)
existing = {
row["name"] for row in conn.execute("PRAGMA table_info(spans)")
}
for name, column_type in _ADDED_COLUMNS:
if name not in existing:
conn.execute(f"ALTER TABLE spans ADD COLUMN {name} {column_type}")
conn.execute(_CREATE_SESSION_INDEX)
conn.execute(_CREATE_TRACE_INDEX)
conn.commit()
Expand Down Expand Up @@ -133,6 +156,56 @@ def _deserialize_attributes(
return {}
return cast(dict[str, AttributeValue], decoded)

def _serialize_events(self, events: Sequence[Event]) -> str:
return json.dumps(
[
{
"name": event.name,
"timestamp": event.timestamp,
"attributes": dict(event.attributes or {}),
}
for event in events
],
ensure_ascii=False,
default=lambda o: "<not serializable>",
)

def _deserialize_events(self, events_json: object) -> list[Event]:
if not isinstance(events_json, (str, bytes, bytearray)):
return []
try:
decoded: object = json.loads(events_json)
except (json.JSONDecodeError, TypeError) as e:
logger.debug("Failed to deserialize span events: %r", e)
return []
if not isinstance(decoded, list):
return []
events: list[Event] = []
for item in decoded:
if not isinstance(item, dict) or not isinstance(item.get("name"), str):
continue
attributes = item.get("attributes")
events.append(
Event(
name=item["name"],
attributes=attributes if isinstance(attributes, dict) else None,
timestamp=item.get("timestamp"),
)
)
return events

def _deserialize_status(
self, status_code: object, status_description: object
) -> Status:
try:
code = StatusCode[str(status_code)]
except KeyError:
return Status(StatusCode.UNSET)
# OpenTelemetry keeps a description only for an error status.
if code is StatusCode.ERROR and isinstance(status_description, str):
return Status(code, status_description)
return Status(code)

def export(self, spans: Sequence[ReadableSpan]) -> SpanExportResult:
try:
with self._lock:
Expand Down Expand Up @@ -167,6 +240,9 @@ def export(self, spans: Sequence[ReadableSpan]) -> SpanExportResult:
session_id,
invocation_id,
self._serialize_attributes(attributes),
span.status.status_code.name,
span.status.description,
self._serialize_events(span.events),
))
conn.executemany(_INSERT_SPAN, rows)
conn.commit()
Expand Down Expand Up @@ -224,6 +300,10 @@ def _row_to_readable_span(self, row: sqlite3.Row) -> ReadableSpan:
attributes=attributes,
start_time=row["start_time_unix_nano"],
end_time=row["end_time_unix_nano"],
status=self._deserialize_status(
row["status_code"], row["status_description"]
),
events=self._deserialize_events(row["events_json"]),
)

def get_all_spans_for_session(self, session_id: str) -> list[ReadableSpan]:
Expand Down
134 changes: 134 additions & 0 deletions tests/unittests/telemetry/test_sqlite_span_exporter.py
Original file line number Diff line number Diff line change
Expand Up @@ -16,11 +16,17 @@

import json
from pathlib import Path
import sqlite3

from google.adk.telemetry.sqlite_span_exporter import SqliteSpanExporter
from opentelemetry.sdk.trace import Event
from opentelemetry.sdk.trace import ReadableSpan
from opentelemetry.sdk.trace import TracerProvider
from opentelemetry.sdk.trace.export import SimpleSpanProcessor
from opentelemetry.sdk.trace.export import SpanExportResult
from opentelemetry.trace import SpanContext
from opentelemetry.trace import Status
from opentelemetry.trace import StatusCode
from opentelemetry.trace import TraceFlags
from opentelemetry.trace import TraceState

Expand Down Expand Up @@ -460,3 +466,131 @@ def test_get_spans_ordered_by_start_time(tmp_path):
assert result[0].context.span_id == 0x100
assert result[1].context.span_id == 0x200
assert result[2].context.span_id == 0x300


def test_round_trip_preserves_error_status_and_events(tmp_path):
exporter = SqliteSpanExporter(db_path=str(tmp_path / "test.db"))
context = SpanContext(
trace_id=0xDEF45,
span_id=0xABC12,
is_remote=False,
trace_flags=TraceFlags(TraceFlags.SAMPLED),
trace_state=TraceState(),
)
span = ReadableSpan(
name="failing_tool",
context=context,
attributes={"gcp.vertex.agent.session_id": "s1"},
start_time=1000,
end_time=2000,
status=Status(StatusCode.ERROR, "ValueError: boom"),
events=[
Event(
name="exception",
attributes={
"exception.type": "ValueError",
"exception.message": "boom",
},
timestamp=1500,
)
],
)

assert exporter.export([span]) == SpanExportResult.SUCCESS
(restored,) = exporter.get_all_spans_for_session("s1")

assert restored.status.status_code is StatusCode.ERROR
assert restored.status.description == "ValueError: boom"
assert len(restored.events) == 1
event = restored.events[0]
assert event.name == "exception"
assert event.timestamp == 1500
assert dict(event.attributes) == {
"exception.type": "ValueError",
"exception.message": "boom",
}


def test_exception_raised_inside_a_real_span_is_read_back(tmp_path):
"""The case from the issue: a run that fails must not read back as clean."""
exporter = SqliteSpanExporter(db_path=str(tmp_path / "test.db"))
provider = TracerProvider()
provider.add_span_processor(SimpleSpanProcessor(exporter))
tracer = provider.get_tracer(__name__)

try:
with tracer.start_as_current_span(
"run", attributes={"gcp.vertex.agent.session_id": "s1"}
):
raise ValueError("boom")
except ValueError:
pass

(restored,) = exporter.get_all_spans_for_session("s1")

assert restored.status.status_code is StatusCode.ERROR
assert "boom" in (restored.status.description or "")
assert [event.name for event in restored.events] == ["exception"]
assert restored.events[0].attributes["exception.type"] == "ValueError"


def test_span_without_status_or_events_reads_back_unset_and_empty(tmp_path):
exporter = SqliteSpanExporter(db_path=str(tmp_path / "test.db"))
exporter.export(
[_create_span(attributes={"gcp.vertex.agent.session_id": "s1"})]
)

(restored,) = exporter.get_all_spans_for_session("s1")

assert restored.status.status_code is StatusCode.UNSET
assert list(restored.events) == []


def test_database_from_older_version_gains_the_new_columns(tmp_path):
"""A file written before status and events were stored keeps working."""
db_path = tmp_path / "old.db"
conn = sqlite3.connect(db_path)
conn.execute("""
CREATE TABLE spans (
span_id TEXT PRIMARY KEY,
trace_id TEXT NOT NULL,
parent_span_id TEXT,
name TEXT NOT NULL,
start_time_unix_nano INTEGER,
end_time_unix_nano INTEGER,
session_id TEXT,
invocation_id TEXT,
attributes_json TEXT
)
""")
conn.execute(
"INSERT INTO spans VALUES (?, ?, ?, ?, ?, ?, ?, ?, ?)",
(
"00000000000abc12",
"0" * 27 + "def45",
None,
"old",
1,
2,
"s1",
None,
"{}",
),
)
conn.commit()
conn.close()

exporter = SqliteSpanExporter(db_path=str(db_path))
exporter.export([
_create_span(
span_id=0x999,
name="new",
attributes={"gcp.vertex.agent.session_id": "s1"},
start_time=5,
)
])

old, new = exporter.get_all_spans_for_session("s1")
assert (old.name, new.name) == ("old", "new")
assert old.status.status_code is StatusCode.UNSET
assert list(old.events) == []
Loading