项目文件夹

文件
2026-07-13 13:22:34 +08:00

894 行
33 KiB
Python

from __future__ import annotations
import copy
import hashlib
import json
import logging
import random
import re
from dataclasses import dataclass
from datetime import datetime, timedelta, timezone
from typing import Any, Literal
import mlflow
from mlflow.demo.base import (
DEMO_EXPERIMENT_NAME,
DEMO_PROMPT_PREFIX,
BaseDemoGenerator,
DemoFeature,
DemoResult,
)
from mlflow.demo.data import (
AGENT_TRACES,
PROMPT_TRACES,
RAG_TRACES,
SESSION_TRACES,
DemoTrace,
MultimodalDemoTrace,
ToolCall,
get_multimodal_traces,
)
from mlflow.entities import SpanType
from mlflow.tracing.constant import SpanAttributeKey, TraceMetadataKey
from mlflow.tracking._tracking_service.utils import _get_store
_logger = logging.getLogger(__name__)
DEMO_VERSION_TAG = "mlflow.demo.version"
DEMO_TRACE_TYPE_TAG = "mlflow.demo.trace_type"
DEMO_SESSION_TURN_TAG = "mlflow.demo.session.turn"
DEMO_START_TIME_TAG = "mlflow.demo.start_time_ms"
DEMO_END_TIME_TAG = "mlflow.demo.end_time_ms"
_TOTAL_TRACES_PER_VERSION = 21
@dataclass(frozen=True)
class _TraceSetResult:
"""Result from generating a set of traces.
Attributes:
trace_ids: List of generated trace IDs.
start_time_ns: Earliest trace start time in nanoseconds.
end_time_ns: Latest trace end time in nanoseconds.
"""
trace_ids: list[str]
start_time_ns: int
end_time_ns: int
def _get_trace_timestamps(trace_index: int, version: str) -> tuple[int, int]:
"""Get deterministic start and end timestamps for a trace.
Distributes traces over the last 7 days with a deterministic pattern
based on the trace index and version. This ensures the demo dashboard
shows activity across the time range.
Args:
trace_index: Index of the trace (0-based) within its version set.
version: "v1" or "v2" - v1 traces are earlier, v2 traces are later.
Returns:
Tuple of (start_time_ns, end_time_ns).
"""
now = datetime.now(timezone.utc)
seven_days_ago = now - timedelta(days=7)
if version == "v1":
day_offset = (trace_index * 3.5) / _TOTAL_TRACES_PER_VERSION
else:
day_offset = 3.5 + (trace_index * 3.5) / _TOTAL_TRACES_PER_VERSION
hash_input = f"{trace_index}:{version}"
hash_val = int(hashlib.md5(hash_input.encode(), usedforsecurity=False).hexdigest()[:8], 16)
hour_offset = (hash_val % 24) / 24
minute_offset = ((hash_val >> 8) % 60) / (60 * 24)
trace_time = seven_days_ago + timedelta(days=day_offset + hour_offset + minute_offset)
duration_ms = 50 + (hash_val % 1950)
start_ns = int(trace_time.timestamp() * 1_000_000_000)
end_ns = start_ns + (duration_ms * 1_000_000)
return start_ns, end_ns
def _estimate_tokens(text: str) -> int:
"""Estimate token count for text (rough approximation: ~4 chars per token)."""
return max(1, len(text) // 4)
@dataclass(frozen=True)
class _Model:
"""Model configuration with name, provider, and pricing."""
name: str
provider: str
pricing: tuple[float, float] # (input $/1M tokens, output $/1M tokens)
# Using three distinct models so the cost breakdown chart shows a nice distribution.
GPT_5_2 = _Model(name="gpt-5.2", provider="openai", pricing=(1.75, 14.00))
CLAUDE_SONNET_4_5 = _Model(name="claude-sonnet-4-5", provider="anthropic", pricing=(3.00, 15.00))
GEMINI_3_PRO = _Model(name="gemini-3-pro", provider="google", pricing=(2.00, 12.00))
_DEMO_MODELS = (GPT_5_2, CLAUDE_SONNET_4_5, GEMINI_3_PRO)
# LLM spans use canonical SDK method names
# Not 100% accurate against production but should be sufficiently understandable for demo purposes
_PROVIDER_TO_LLM_SPAN_NAME = {
"openai": "chat.completions.create",
"anthropic": "messages.create",
"google": "generate_content",
}
def _compute_cost(model: _Model, prompt_tokens: int, completion_tokens: int) -> dict[str, float]:
"""Compute synthetic cost using approximate per-model pricing."""
input_rate, output_rate = model.pricing
input_cost = prompt_tokens * input_rate / 1_000_000
output_cost = completion_tokens * output_rate / 1_000_000
return {
"input_cost": input_cost,
"output_cost": output_cost,
"total_cost": input_cost + output_cost,
}
def _json_type(value: Any) -> str:
# Intentionally shallow: nested dicts/lists are reported as bare "object"/"array"
# without `properties`/`items`. Fine for a demo schema where we only need the
# top-level parameter shape; not a general-purpose JSON Schema generator.
if isinstance(value, bool):
return "boolean"
if isinstance(value, int):
return "integer"
if isinstance(value, float):
return "number"
if isinstance(value, list):
return "array"
if isinstance(value, dict):
return "object"
return "string"
def _tool_schemas(tools: list[ToolCall]) -> list[dict[str, Any]]:
"""
Build OpenAI-style function schemas from a list of ToolCall objects.
Referenced from: https://developers.openai.com/api/docs/guides/function-calling
"""
return [
{
"type": "function",
"function": {
"name": tool.name,
"description": f"Call the {tool.name} tool.",
"parameters": {
"type": "object",
"properties": {k: {"type": _json_type(v)} for k, v in tool.input.items()},
"required": list(tool.input.keys()),
},
},
}
for tool in tools
]
def _llm_attributes(model: _Model, in_toks: int, out_toks: int) -> dict[str, Any]:
return {
SpanAttributeKey.CHAT_USAGE: {
"input_tokens": in_toks,
"output_tokens": out_toks,
"total_tokens": in_toks + out_toks,
},
SpanAttributeKey.MODEL: model.name,
SpanAttributeKey.MODEL_PROVIDER: model.provider,
SpanAttributeKey.LLM_COST: _compute_cost(model, in_toks, out_toks),
}
def _emit_react_children(
root,
tools: list[ToolCall],
model: _Model,
system_content: str,
user_query: str,
response: str,
start_ns: int,
end_ns: int,
prior_messages: list[dict[str, Any]] | None = None,
) -> None:
"""Emit ReAct-style child spans under `root`.
For N tools, emits N+1 LLM spans alternating with N TOOL spans:
LLM(decide call_1) → TOOL(1) → LLM(decide call_2) → TOOL(2) → … → LLM(final).
If there are no tools, emits a single LLM span.
`prior_messages` is the running conversation history from earlier turns in the
same session. It is inserted between the system prompt and the current user query
so the LLM sees the full context, the way a real stateful chat agent would.
"""
span_name = _PROVIDER_TO_LLM_SPAN_NAME[model.provider]
tool_schemas = _tool_schemas(tools)
schemas_token_overhead = _estimate_tokens(json.dumps(tool_schemas))
messages = [{"role": "system", "content": system_content}]
messages.extend(prior_messages or [])
messages.append({"role": "user", "content": user_query})
# Each span gets a jittered duration so per-span latency varies trace-to-trace.
# Pre-compute all per-span durations and rescale them to fit exactly into the
# `[start_ns + 5_000, end_ns - 5_000]` window — this guarantees spans stay
# contiguous and non-overlapping even when high jitter draws would otherwise
# push the cursor past the end. Seeded by start_ns for determinism.
total_spans = 2 * len(tools) + 1
budget = max(total_spans, end_ns - start_ns - 10_000)
rng = random.Random(start_ns)
# Each span's raw weight is uniformly drawn from [0.2, 1.8], giving the longest
# span in a trace up to ~9x the duration of the shortest (1.8 / 0.2). The mean
# of 1.0 keeps the expected sum equal to `total_spans`, so after rescaling
# below each span occupies roughly its drawn fraction of the budget. Tweak this
# range to widen or narrow the visible latency spread in the timeline.
raw_durations = [rng.uniform(0.2, 1.8) for _ in range(total_spans)]
total_raw = sum(raw_durations)
span_durations = [max(1, int(d / total_raw * budget)) for d in raw_durations]
cursor = start_ns + 5_000
for idx, tool in enumerate(tools, start=1):
call_id = f"call_{idx:03d}"
arguments_json = json.dumps(tool.input)
tool_call = {
"id": call_id,
"type": "function",
"function": {"name": tool.name, "arguments": arguments_json},
}
in_toks = _estimate_tokens(json.dumps(messages)) + schemas_token_overhead
out_toks = _estimate_tokens(tool.name + arguments_json) + 5
llm = mlflow.start_span_no_context(
name=span_name,
span_type=SpanType.LLM,
parent_span=root,
inputs={"messages": list(messages), "model": model.name, "tools": tool_schemas},
attributes=_llm_attributes(model, in_toks, out_toks),
start_time_ns=cursor,
)
llm.set_outputs({
"choices": [
{
"message": {
"role": "assistant",
"content": None,
"tool_calls": [tool_call],
},
"finish_reason": "tool_calls",
}
]
})
cursor += span_durations[2 * (idx - 1)]
llm.end(end_time_ns=cursor)
messages.append({"role": "assistant", "content": None, "tool_calls": [tool_call]})
tool_span = mlflow.start_span_no_context(
name=tool.name,
span_type=SpanType.TOOL,
parent_span=root,
inputs=tool.input,
start_time_ns=cursor,
)
tool_span.set_outputs(tool.output)
cursor += span_durations[2 * (idx - 1) + 1]
tool_span.end(end_time_ns=cursor)
messages.append({
"role": "tool",
"tool_call_id": call_id,
"content": json.dumps(tool.output),
})
in_toks = _estimate_tokens(json.dumps(messages)) + schemas_token_overhead
out_toks = _estimate_tokens(response)
final = mlflow.start_span_no_context(
name=span_name,
span_type=SpanType.LLM,
parent_span=root,
inputs={"messages": list(messages), "model": model.name, "tools": tool_schemas},
attributes=_llm_attributes(model, in_toks, out_toks),
start_time_ns=cursor,
)
final.set_outputs({"choices": [{"message": {"role": "assistant", "content": response}}]})
final.end(end_time_ns=end_ns - 5_000)
class TracesDemoGenerator(BaseDemoGenerator):
"""Generates demo traces for the MLflow UI.
Creates two sets of traces showing agent improvement:
- V1 traces: Initial/baseline agent (uses v1_response)
- V2 traces: Improved agent after updates (uses v2_response)
Both versions use the same inputs but produce different outputs,
simulating an agent improvement workflow.
Trace types generated:
- RAG: Document retrieval and generation pipeline
- Agent: Tool-using agent with function calls
- Prompt: Prompt template-based generation
- Session: Multi-turn conversation sessions
"""
name = DemoFeature.TRACES
version = 3
def generate(self) -> DemoResult:
self._restore_experiment_if_deleted()
experiment = mlflow.set_experiment(DEMO_EXPERIMENT_NAME)
mlflow.MlflowClient().set_experiment_tag(
experiment.experiment_id, "mlflow.experimentKind", "genai_development"
)
mlflow.set_experiment_tag(
"mlflow.note.content",
"Sample experiment with pre-populated demo data including traces, evaluations, "
"and prompts. Explore MLflow's GenAI features with this experiment.",
)
v1_result = self._generate_trace_set("v1")
v2_result = self._generate_trace_set("v2")
all_trace_ids = v1_result.trace_ids + v2_result.trace_ids
# Store the overall time range of demo data as experiment tags
overall_start_ms = min(v1_result.start_time_ns, v2_result.start_time_ns) // 1_000_000
overall_end_ms = max(v1_result.end_time_ns, v2_result.end_time_ns) // 1_000_000
mlflow.set_experiment_tag(DEMO_START_TIME_TAG, str(overall_start_ms))
mlflow.set_experiment_tag(DEMO_END_TIME_TAG, str(overall_end_ms))
return DemoResult(
feature=self.name,
entity_ids=all_trace_ids,
navigation_url=f"#/experiments/{experiment.experiment_id}",
)
def _generate_trace_set(self, version: Literal["v1", "v2"]) -> _TraceSetResult:
"""Generate a complete set of traces for the given version."""
trace_ids = []
trace_index = 0
min_start_ns = float("inf")
max_end_ns = 0
for trace_def in RAG_TRACES:
start_ns, end_ns = _get_trace_timestamps(trace_index, version)
min_start_ns = min(min_start_ns, start_ns)
max_end_ns = max(max_end_ns, end_ns)
if trace_id := self._create_rag_trace(trace_def, version, start_ns, end_ns):
trace_ids.append(trace_id)
trace_index += 1
for trace_def in AGENT_TRACES:
start_ns, end_ns = _get_trace_timestamps(trace_index, version)
min_start_ns = min(min_start_ns, start_ns)
max_end_ns = max(max_end_ns, end_ns)
if trace_id := self._create_agent_trace(trace_def, version, start_ns, end_ns):
trace_ids.append(trace_id)
trace_index += 1
for idx, trace_def in enumerate(PROMPT_TRACES):
start_ns, end_ns = _get_trace_timestamps(trace_index, version)
min_start_ns = min(min_start_ns, start_ns)
max_end_ns = max(max_end_ns, end_ns)
prompt_version_num = str(idx % 2 + 1) if version == "v1" else str(idx % 2 + 3)
if trace_id := self._create_prompt_trace(
trace_def, version, start_ns, end_ns, prompt_version_num
):
trace_ids.append(trace_id)
trace_index += 1
for trace_def in get_multimodal_traces():
start_ns, end_ns = _get_trace_timestamps(trace_index, version)
min_start_ns = min(min_start_ns, start_ns)
max_end_ns = max(max_end_ns, end_ns)
if trace_id := self._create_multimodal_trace(trace_def, version, start_ns, end_ns):
trace_ids.append(trace_id)
trace_index += 1
session_result = self._create_session_traces(version, trace_index)
trace_ids.extend(session_result.trace_ids)
min_start_ns = min(min_start_ns, session_result.start_time_ns)
max_end_ns = max(max_end_ns, session_result.end_time_ns)
return _TraceSetResult(
trace_ids=trace_ids,
start_time_ns=int(min_start_ns),
end_time_ns=int(max_end_ns),
)
def _data_exists(self) -> bool:
store = _get_store()
try:
experiment = store.get_experiment_by_name(DEMO_EXPERIMENT_NAME)
if experiment is None or experiment.lifecycle_stage != "active":
return False
traces = mlflow.search_traces(
locations=[experiment.experiment_id],
max_results=1,
flush=True,
)
return len(traces) > 0
except Exception:
_logger.debug("Failed to check if demo data exists", exc_info=True)
return False
def delete_demo(self) -> None:
store = _get_store()
try:
experiment = store.get_experiment_by_name(DEMO_EXPERIMENT_NAME)
if experiment is None:
return
client = mlflow.MlflowClient()
traces = client.search_traces(
locations=[experiment.experiment_id],
max_results=200,
)
if trace_ids := [trace.info.trace_id for trace in traces]:
try:
client.delete_traces(
experiment_id=experiment.experiment_id,
trace_ids=trace_ids,
)
except Exception:
pass
except Exception:
_logger.debug("Failed to delete demo traces", exc_info=True)
def _restore_experiment_if_deleted(self) -> None:
"""Restore the demo experiment if it was soft-deleted."""
store = _get_store()
try:
experiment = store.get_experiment_by_name(DEMO_EXPERIMENT_NAME)
if experiment is not None and experiment.lifecycle_stage == "deleted":
_logger.info("Restoring soft-deleted demo experiment")
client = mlflow.MlflowClient()
client.restore_experiment(experiment.experiment_id)
except Exception:
_logger.debug("Failed to check/restore demo experiment", exc_info=True)
def _get_response(self, trace_def: DemoTrace, version: Literal["v1", "v2"]) -> str:
"""Get the appropriate response based on version."""
return trace_def.v1_response if version == "v1" else trace_def.v2_response
def _create_rag_trace(
self,
trace_def: DemoTrace,
version: Literal["v1", "v2"],
start_ns: int,
end_ns: int,
) -> str | None:
"""Create a RAG pipeline trace: embed -> retrieve -> generate."""
response = self._get_response(trace_def, version)
prompt_tokens = _estimate_tokens(trace_def.query) + 50
completion_tokens = _estimate_tokens(response)
total_duration = end_ns - start_ns
embed_end = start_ns + int(total_duration * 0.1)
retrieve_end = embed_end + int(total_duration * 0.2)
llm_start = retrieve_end
llm_end = end_ns - int(total_duration * 0.05)
root = mlflow.start_span_no_context(
name="rag_pipeline",
span_type=SpanType.CHAIN,
inputs={"messages": [{"role": "user", "content": trace_def.query}]},
metadata={DEMO_VERSION_TAG: version, DEMO_TRACE_TYPE_TAG: "rag"},
start_time_ns=start_ns,
)
embed = mlflow.start_span_no_context(
name="embed_query",
span_type=SpanType.EMBEDDING,
parent_span=root,
inputs={"text": trace_def.query},
start_time_ns=start_ns + 1000,
)
embedding = [random.uniform(-1, 1) for _ in range(384)]
embed.set_outputs({"embedding": embedding[:5], "dimensions": 384})
embed.end(end_time_ns=embed_end)
retrieve = mlflow.start_span_no_context(
name="retrieve_docs",
span_type=SpanType.RETRIEVER,
parent_span=root,
inputs={"embedding": embedding[:5], "top_k": 3},
start_time_ns=embed_end + 1000,
)
docs = [
{"id": f"doc_{i}", "score": round(0.7 + random.uniform(0, 0.25), 2)} for i in range(3)
]
retrieve.set_outputs({"documents": docs})
retrieve.end(end_time_ns=retrieve_end)
model = GPT_5_2
llm = mlflow.start_span_no_context(
name=_PROVIDER_TO_LLM_SPAN_NAME[model.provider],
span_type=SpanType.LLM,
parent_span=root,
inputs={
"messages": [
{"role": "system", "content": "You are an MLflow assistant."},
{"role": "user", "content": trace_def.query},
],
"context": docs,
"model": model.name,
},
attributes={
SpanAttributeKey.CHAT_USAGE: {
"input_tokens": prompt_tokens,
"output_tokens": completion_tokens,
"total_tokens": prompt_tokens + completion_tokens,
},
SpanAttributeKey.MODEL: model.name,
SpanAttributeKey.MODEL_PROVIDER: model.provider,
SpanAttributeKey.LLM_COST: _compute_cost(model, prompt_tokens, completion_tokens),
},
start_time_ns=llm_start,
)
llm.set_outputs({"choices": [{"message": {"role": "assistant", "content": response}}]})
llm.end(end_time_ns=llm_end)
root.set_outputs({"choices": [{"message": {"role": "assistant", "content": response}}]})
root.end(end_time_ns=end_ns)
return root.trace_id
def _create_agent_trace(
self,
trace_def: DemoTrace,
version: Literal["v1", "v2"],
start_ns: int,
end_ns: int,
) -> str | None:
response = self._get_response(trace_def, version)
root = mlflow.start_span_no_context(
name="agent",
span_type=SpanType.AGENT,
inputs={"messages": [{"role": "user", "content": trace_def.query}]},
metadata={DEMO_VERSION_TAG: version, DEMO_TRACE_TYPE_TAG: "agent"},
start_time_ns=start_ns,
)
_emit_react_children(
root=root,
tools=trace_def.tools,
model=CLAUDE_SONNET_4_5,
system_content="You are a helpful assistant with tools.",
user_query=trace_def.query,
response=response,
start_ns=start_ns,
end_ns=end_ns,
)
root.set_outputs({"choices": [{"message": {"role": "assistant", "content": response}}]})
root.end(end_time_ns=end_ns)
return root.trace_id
def _create_prompt_trace(
self,
trace_def: DemoTrace,
version: Literal["v1", "v2"],
start_ns: int,
end_ns: int,
prompt_version: str = "1",
) -> str | None:
"""Create a prompt-based trace showing template rendering and generation.
Fetches the actual registered prompt template and renders it with appropriate
variables to ensure trace contents match the linked prompt version.
"""
response = self._get_response(trace_def, version)
if trace_def.prompt_template is None:
return None
full_prompt_name = f"{DEMO_PROMPT_PREFIX}.prompts.{trace_def.prompt_template.prompt_name}"
try:
client = mlflow.MlflowClient()
prompt_version_obj = client.get_prompt_version(
name=full_prompt_name,
version=prompt_version,
)
actual_template = prompt_version_obj.template
except Exception:
actual_template = trace_def.prompt_template.template
variables = self._get_prompt_variables(
trace_def.prompt_template.prompt_name,
trace_def.query,
trace_def.prompt_template.variables,
)
rendered_prompt = self._render_template(actual_template, variables)
prompt_tokens = _estimate_tokens(rendered_prompt) + 20
completion_tokens = _estimate_tokens(response)
total_duration = end_ns - start_ns
render_end = start_ns + int(total_duration * 0.1)
llm_start = render_end + 1000
root = mlflow.start_span_no_context(
name="prompt_chain",
span_type=SpanType.CHAIN,
inputs={
"messages": [{"role": "user", "content": trace_def.query}],
"template_variables": variables,
},
metadata={DEMO_VERSION_TAG: version, DEMO_TRACE_TYPE_TAG: "prompt"},
start_time_ns=start_ns,
)
render = mlflow.start_span_no_context(
name="render_prompt",
span_type=SpanType.CHAIN,
parent_span=root,
inputs={
"template": actual_template,
"template_variables": variables,
},
start_time_ns=start_ns + 1000,
)
render.set_outputs({"rendered_prompt": rendered_prompt})
render.end(end_time_ns=render_end)
model = GEMINI_3_PRO
llm = mlflow.start_span_no_context(
name=_PROVIDER_TO_LLM_SPAN_NAME[model.provider],
span_type=SpanType.LLM,
parent_span=root,
inputs={
"messages": [
{"role": "user", "content": rendered_prompt},
],
"model": model.name,
},
attributes={
SpanAttributeKey.CHAT_USAGE: {
"input_tokens": prompt_tokens,
"output_tokens": completion_tokens,
"total_tokens": prompt_tokens + completion_tokens,
},
SpanAttributeKey.MODEL: model.name,
SpanAttributeKey.MODEL_PROVIDER: model.provider,
SpanAttributeKey.LLM_COST: _compute_cost(model, prompt_tokens, completion_tokens),
},
start_time_ns=llm_start,
)
llm.set_outputs({"choices": [{"message": {"role": "assistant", "content": response}}]})
llm.end(end_time_ns=end_ns - 5000)
root.set_outputs({"choices": [{"message": {"role": "assistant", "content": response}}]})
root.end(end_time_ns=end_ns)
trace_id = root.trace_id
self._link_prompt_to_trace(trace_def.prompt_template.prompt_name, trace_id, prompt_version)
return trace_id
def _create_multimodal_trace(
self,
trace_def: MultimodalDemoTrace,
version: Literal["v1", "v2"],
start_ns: int,
end_ns: int,
) -> str | None:
"""Create a multimodal trace with pre-built inputs/outputs."""
response_text = (
trace_def.v1_response_text if version == "v1" else trace_def.v2_response_text
)
prompt_tokens = 200
completion_tokens = _estimate_tokens(response_text)
model = GPT_5_2
# Deep copy to avoid mutating shared trace definition data
outputs = copy.deepcopy(trace_def.outputs)
# Inject version-specific response text into outputs
match outputs:
case {"choices": [*choices]}:
for choice in choices:
match choice:
case {"message": {"content": None, **rest}} if "audio" not in rest:
choice["message"]["content"] = response_text
root = mlflow.start_span_no_context(
name=trace_def.name,
span_type=trace_def.span_type,
inputs=trace_def.inputs,
attributes={
SpanAttributeKey.MESSAGE_FORMAT: "openai",
SpanAttributeKey.CHAT_USAGE: {
"input_tokens": prompt_tokens,
"output_tokens": completion_tokens,
"total_tokens": prompt_tokens + completion_tokens,
},
SpanAttributeKey.MODEL: model.name,
SpanAttributeKey.MODEL_PROVIDER: model.provider,
SpanAttributeKey.LLM_COST: _compute_cost(model, prompt_tokens, completion_tokens),
},
metadata={DEMO_VERSION_TAG: version, DEMO_TRACE_TYPE_TAG: "multimodal"},
start_time_ns=start_ns,
)
root.set_outputs(outputs)
root.end(end_time_ns=end_ns)
return root.trace_id
def _link_prompt_to_trace(
self, short_prompt_name: str, trace_id: str, prompt_version: str = "1"
) -> None:
full_prompt_name = f"{DEMO_PROMPT_PREFIX}.prompts.{short_prompt_name}"
try:
client = mlflow.MlflowClient()
prompt_version_obj = client.get_prompt_version(
name=full_prompt_name,
version=prompt_version,
)
client.link_prompt_versions_to_trace(
prompt_versions=[prompt_version_obj],
trace_id=trace_id,
)
except Exception:
_logger.debug(
"Failed to link prompt %s v%s to trace %s",
full_prompt_name,
prompt_version,
trace_id,
exc_info=True,
)
def _get_prompt_variables(
self, prompt_name: str, query: str, base_variables: dict[str, str]
) -> dict[str, str]:
"""Get complete variable set for a prompt type.
Combines base variables from the trace definition with additional
variables that may be needed for more advanced prompt versions.
"""
variables = dict(base_variables)
if "query" not in variables:
variables["query"] = query
if prompt_name == "customer-support":
variables.setdefault("company_name", "TechCorp")
variables.setdefault("context", "Customer has been with us for 2 years, premium tier.")
elif prompt_name == "document-summarizer":
variables.setdefault("max_words", "150")
variables.setdefault("audience", "technical professionals")
variables.setdefault(
"document",
variables.get("query", "Sample document content for summarization."),
)
elif prompt_name == "code-reviewer":
variables.setdefault("language", "python")
variables.setdefault("focus_areas", "security, performance, readability")
variables.setdefault("severity_levels", "critical, warning, suggestion")
variables.setdefault("code", variables.get("query", "def example(): pass"))
return variables
def _render_template(
self, template: str | list[dict[str, str]], variables: dict[str, str]
) -> str:
"""Render a prompt template with variables.
Handles both string templates and chat-format templates (list of messages).
"""
def substitute(text: str, vars_dict: dict[str, str]) -> str:
for key, value in vars_dict.items():
text = re.sub(r"\{\{\s*" + key + r"\s*\}\}", str(value), text)
return text
if isinstance(template, str):
return substitute(template, variables)
elif isinstance(template, list):
rendered_parts = []
for msg in template:
role = msg.get("role", "user")
content = substitute(msg.get("content", ""), variables)
rendered_parts.append(f"[{role}]: {content}")
return "\n\n".join(rendered_parts)
else:
return str(template)
def _create_session_traces(
self, version: Literal["v1", "v2"], start_index: int
) -> _TraceSetResult:
"""Create multi-turn conversation session traces."""
trace_ids = []
current_session = None
turn_counter = 0
trace_index = start_index
min_start_ns = float("inf")
max_end_ns = 0
prior_by_session: dict[str, list[dict[str, Any]]] = {}
for trace_def in SESSION_TRACES:
if trace_def.session_id != current_session:
current_session = trace_def.session_id
turn_counter = 0
turn_counter += 1
versioned_session_id = f"{trace_def.session_id}-{version}"
start_ns, end_ns = _get_trace_timestamps(trace_index, version)
min_start_ns = min(min_start_ns, start_ns)
max_end_ns = max(max_end_ns, end_ns)
prior = prior_by_session.setdefault(versioned_session_id, [])
if trace_id := self._create_session_turn_trace(
trace_def,
turn_counter,
version,
versioned_session_id,
start_ns,
end_ns,
prior_messages=prior,
):
trace_ids.append(trace_id)
prior.append({"role": "user", "content": trace_def.query})
prior.append({"role": "assistant", "content": self._get_response(trace_def, version)})
trace_index += 1
return _TraceSetResult(
trace_ids=trace_ids,
start_time_ns=int(min_start_ns),
end_time_ns=int(max_end_ns),
)
def _create_session_turn_trace(
self,
trace_def: DemoTrace,
turn: int,
version: Literal["v1", "v2"],
versioned_session_id: str,
start_ns: int,
end_ns: int,
prior_messages: list[dict[str, Any]] | None = None,
) -> str | None:
"""Create a single turn in a conversation session."""
response = self._get_response(trace_def, version)
root = mlflow.start_span_no_context(
name="chat_agent",
span_type=SpanType.AGENT,
inputs={"messages": [{"role": "user", "content": trace_def.query}]},
metadata={
TraceMetadataKey.TRACE_SESSION: versioned_session_id,
TraceMetadataKey.TRACE_USER: trace_def.session_user or "user",
DEMO_VERSION_TAG: version,
DEMO_TRACE_TYPE_TAG: "session",
DEMO_SESSION_TURN_TAG: str(turn),
},
start_time_ns=start_ns,
)
_emit_react_children(
root=root,
tools=trace_def.tools,
model=_DEMO_MODELS[turn % len(_DEMO_MODELS)],
system_content="You are an MLflow assistant.",
user_query=trace_def.query,
response=response,
start_ns=start_ns,
end_ns=end_ns,
prior_messages=prior_messages,
)
root.set_outputs({"choices": [{"message": {"role": "assistant", "content": response}}]})
root.end(end_time_ns=end_ns)
return root.trace_id