Skip to content

Commit 21f47e6

Browse files
Copilothiyouga
andauthored
Refactor created_at stamping and tracing into streaming_response
Agent-Logs-Url: https://github.com/Prism-Shadow/AgentHub/sessions/57d3d757-2462-486d-b417-46e62dbe6cb8 Co-authored-by: hiyouga <16256802+hiyouga@users.noreply.github.com>
1 parent 044a51d commit 21f47e6

2 files changed

Lines changed: 48 additions & 34 deletions

File tree

src_py/agenthub/base_client.py

Lines changed: 20 additions & 17 deletions
Original file line numberDiff line numberDiff line change
@@ -13,7 +13,7 @@
1313
# limitations under the License.
1414

1515
from abc import ABC, abstractmethod
16-
from typing import Any, AsyncIterator, cast
16+
from typing import Any, AsyncIterator
1717
import time
1818

1919
from .types import ContentItem, FinishReason, UniConfig, UniEvent, UniMessage, UsageMetadata
@@ -123,15 +123,13 @@ def concat_uni_events_to_uni_message(self, events: list[UniEvent]) -> UniMessage
123123
finish_reason = event.get("finish_reason") # finish_reason is taken from the last event
124124
created_at = event.get("created_at") # created_at is taken from the last event
125125

126-
result: UniMessage = {
126+
return {
127127
"role": "assistant",
128128
"content_items": content_items,
129129
"usage_metadata": usage_metadata,
130130
"finish_reason": finish_reason,
131+
"created_at": created_at,
131132
}
132-
if created_at is not None:
133-
result["created_at"] = created_at
134-
return result
135133

136134
@abstractmethod
137135
async def _streaming_response_internal(
@@ -173,14 +171,29 @@ async def streaming_response(
173171
Yields:
174172
Universal events from the streaming response
175173
"""
174+
# Stamp any messages that don't yet have a created_at timestamp
175+
for msg in messages:
176+
if "created_at" not in msg:
177+
msg["created_at"] = int(time.time() * 1000)
178+
176179
last_event: UniEvent | None = None
180+
events = []
177181
async for event in self._streaming_response_internal(messages, config):
178182
event["created_at"] = int(time.time() * 1000)
179183
last_event = event
184+
events.append(event)
180185
yield event
181186

182187
self._validate_last_event(last_event)
183188

189+
# Save history to file if trace_id is specified
190+
if config.get("trace_id") and events:
191+
from .integration.tracer import Tracer
192+
193+
assistant_message = self.concat_uni_events_to_uni_message(events)
194+
tracer = Tracer()
195+
tracer.save_history(self._model, messages + [assistant_message], config["trace_id"], config)
196+
184197
async def streaming_response_stateful(
185198
self,
186199
message: UniMessage,
@@ -200,10 +213,6 @@ async def streaming_response_stateful(
200213
Yields:
201214
Universal events from the streaming response
202215
"""
203-
# Stamp input message with current time if not provided
204-
if "created_at" not in message:
205-
message = cast(UniMessage, {**message, "created_at": int(time.time() * 1000)})
206-
207216
# Build a temporary messages list for inference without mutating history yet
208217
temp_messages = self._history + [message]
209218

@@ -214,18 +223,12 @@ async def streaming_response_stateful(
214223
yield event
215224

216225
# Only update history after successful inference
226+
# temp_messages[-1] is the user message, now stamped with created_at by streaming_response
217227
if events:
218228
assistant_message = self.concat_uni_events_to_uni_message(events)
219-
self._history.append(message)
229+
self._history.append(temp_messages[-1])
220230
self._history.append(assistant_message)
221231

222-
# Save history to file if trace_id is specified
223-
if config.get("trace_id"):
224-
from .integration.tracer import Tracer
225-
226-
tracer = Tracer()
227-
tracer.save_history(self._model, self._history, config["trace_id"], config)
228-
229232
@staticmethod
230233
def _validate_last_event(last_event: UniEvent | None) -> None:
231234
"""Validate that the last event has usage_metadata and finish_reason.

src_ts/src/baseClient.ts

Lines changed: 28 additions & 17 deletions
Original file line numberDiff line numberDiff line change
@@ -127,10 +127,8 @@ export abstract class LLMClient {
127127
content_items: contentItems,
128128
usage_metadata: usageMetadata,
129129
finish_reason: finishReason,
130+
created_at: createdAt ?? undefined,
130131
};
131-
if (createdAt !== null) {
132-
result.created_at = createdAt;
133-
}
134132
return result;
135133
}
136134

@@ -162,13 +160,37 @@ export abstract class LLMClient {
162160
messages: UniMessage[];
163161
config: UniConfig;
164162
}): AsyncGenerator<UniEvent> {
163+
const { messages, config } = options;
164+
165+
// Stamp any messages that don't yet have a created_at timestamp
166+
for (const msg of messages) {
167+
if (msg.created_at == null) {
168+
msg.created_at = Date.now();
169+
}
170+
}
171+
165172
let lastEvent: UniEvent | null = null;
173+
const events: UniEvent[] = [];
166174
for await (const event of this._streamingResponseInternal(options)) {
167175
event.created_at = Date.now();
168176
lastEvent = event;
177+
events.push(event);
169178
yield event;
170179
}
171180
LLMClient._validateLastEvent(lastEvent);
181+
182+
// Save history to file if trace_id is specified
183+
if (config.trace_id && events.length > 0) {
184+
const { Tracer } = await import("./integration/tracer");
185+
const assistantMessage = this.concatUniEventsToUniMessage(events);
186+
const tracer = new Tracer();
187+
tracer.saveHistory(
188+
this._model,
189+
[...messages, assistantMessage],
190+
config.trace_id,
191+
config,
192+
);
193+
}
172194
}
173195

174196
/**
@@ -186,13 +208,7 @@ export abstract class LLMClient {
186208
message: UniMessage;
187209
config: UniConfig;
188210
}): AsyncGenerator<UniEvent> {
189-
let { message } = options;
190-
const { config } = options;
191-
192-
// Stamp input message with current time if not provided (shallow copy to avoid mutating caller's object)
193-
if (message.created_at == null) {
194-
message = { ...message, created_at: Date.now() };
195-
}
211+
const { message, config } = options;
196212

197213
const tempMessages = [...this._history, message];
198214

@@ -205,17 +221,12 @@ export abstract class LLMClient {
205221
yield event;
206222
}
207223

224+
// tempMessages[-1] is the user message, now stamped with created_at by streamingResponse
208225
if (events.length > 0) {
209226
const assistantMessage = this.concatUniEventsToUniMessage(events);
210-
this._history.push(message);
227+
this._history.push(tempMessages[tempMessages.length - 1]);
211228
this._history.push(assistantMessage);
212229
}
213-
214-
if (config.trace_id) {
215-
const { Tracer } = await import("./integration/tracer");
216-
const tracer = new Tracer();
217-
tracer.saveHistory(this._model, this._history, config.trace_id, config);
218-
}
219230
}
220231

221232
/**

0 commit comments

Comments
 (0)