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
2 changes: 2 additions & 0 deletions components/model/agenticopenai/consts.go
Original file line number Diff line number Diff line change
Expand Up @@ -22,6 +22,8 @@ const responsesImplType = "AgenticOpenAI/Responses"

const defaultBaseURL = "https://api.openai.com/v1"

const keyOfCacheWriteTokens = "_eino_openai_cache_write_tokens"

type ServerToolName string

const (
Expand Down
58 changes: 58 additions & 0 deletions components/model/agenticopenai/message_extra.go
Original file line number Diff line number Diff line change
@@ -0,0 +1,58 @@
/*
* Copyright 2026 CloudWeGo Authors
*
* 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
*
* https://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.
*/

package agenticopenai

import (
"strconv"

"github.com/cloudwego/eino/schema"
"github.com/openai/openai-go/v3/responses"
)

// GetCacheWriteTokens returns the OpenAI cache_write_tokens count from the
// message. This is the number of input tokens written to the prompt cache.
// Pricing for cache writes depends on the model.
//
// The cache-read side is available through the standard token usage path:
//
// msg.ResponseMeta.TokenUsage.PromptTokenDetails.CachedTokens
//
// When streaming, OpenAI reports cache-write usage on a response lifecycle
// event. schema.ConcatAgenticMessages merges Extra maps with a last-value-wins
// policy for int, so the final concatenated message preserves the count.
func GetCacheWriteTokens(msg *schema.AgenticMessage) (int, bool) {
if msg == nil || msg.Extra == nil {
return 0, false
}
tokens, ok := msg.Extra[keyOfCacheWriteTokens].(int)
return tokens, ok
}

func cacheWriteTokensExtra(resp *responses.Response) map[string]any {
if resp == nil {
return nil
}
field, ok := resp.Usage.InputTokensDetails.JSON.ExtraFields["cache_write_tokens"]
if !ok {
return nil
}
tokens, err := strconv.Atoi(field.Raw())
if err != nil || tokens <= 0 {
return nil
}
return map[string]any{keyOfCacheWriteTokens: tokens}
}
175 changes: 175 additions & 0 deletions components/model/agenticopenai/message_extra_test.go
Original file line number Diff line number Diff line change
@@ -0,0 +1,175 @@
/*
* Copyright 2026 CloudWeGo Authors
*
* 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
*
* https://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.
*/

package agenticopenai

import (
"encoding/json"
"fmt"
"testing"

"github.com/cloudwego/eino/components/model"
"github.com/cloudwego/eino/schema"
"github.com/openai/openai-go/v3/responses"
)

func TestGetCacheWriteTokens(t *testing.T) {
tests := []struct {
name string
msg *schema.AgenticMessage
wantTokens int
wantOK bool
}{
{name: "nil message"},
{name: "nil extra", msg: &schema.AgenticMessage{}},
{name: "missing key", msg: &schema.AgenticMessage{Extra: map[string]any{"other": 42}}},
{
name: "cache write tokens present",
msg: &schema.AgenticMessage{Extra: map[string]any{
keyOfCacheWriteTokens: 1234,
}},
wantTokens: 1234,
wantOK: true,
},
}

for _, tt := range tests {
t.Run(tt.name, func(t *testing.T) {
gotTokens, gotOK := GetCacheWriteTokens(tt.msg)
if gotTokens != tt.wantTokens || gotOK != tt.wantOK {
t.Fatalf("GetCacheWriteTokens() = (%d, %v), want (%d, %v)",
gotTokens, gotOK, tt.wantTokens, tt.wantOK)
}
})
}
}

func TestCacheWriteTokensSetOnGeneratedMessage(t *testing.T) {
t.Run("zero cache write tokens are omitted", func(t *testing.T) {
resp := mustUnmarshalResponse(t, `{
"id": "resp_1",
"status": "completed",
"output": [],
"usage": {
"input_tokens": 100,
"input_tokens_details": {"cached_tokens": 0, "cache_write_tokens": 0},
"output_tokens": 20,
"output_tokens_details": {"reasoning_tokens": 0},
"total_tokens": 120
}
}`)

msg, err := toOutputMessage(resp, &model.Options{})
if err != nil {
t.Fatal(err)
}
if tokens, ok := GetCacheWriteTokens(msg); ok || tokens != 0 {
t.Fatalf("expected (0, false), got (%d, %v)", tokens, ok)
}
})

t.Run("nonzero cache write tokens are exposed", func(t *testing.T) {
resp := mustUnmarshalResponse(t, `{
"id": "resp_1",
"status": "completed",
"output": [],
"usage": {
"input_tokens": 600,
"input_tokens_details": {"cached_tokens": 100, "cache_write_tokens": 500},
"output_tokens": 20,
"output_tokens_details": {"reasoning_tokens": 0},
"total_tokens": 620
}
}`)

msg, err := toOutputMessage(resp, &model.Options{})
if err != nil {
t.Fatal(err)
}
if tokens, ok := GetCacheWriteTokens(msg); !ok || tokens != 500 {
t.Fatalf("expected (500, true), got (%d, %v)", tokens, ok)
}
})
}

func TestCacheWriteTokensInvalidValuesAreOmitted(t *testing.T) {
tests := []struct {
name string
cacheWriteDetails string
}{
{name: "absent", cacheWriteDetails: `"cached_tokens": 0`},
{name: "null", cacheWriteDetails: `"cached_tokens": 0, "cache_write_tokens": null`},
{name: "string", cacheWriteDetails: `"cached_tokens": 0, "cache_write_tokens": "300"`},
{name: "fractional", cacheWriteDetails: `"cached_tokens": 0, "cache_write_tokens": 1.5`},
{name: "negative", cacheWriteDetails: `"cached_tokens": 0, "cache_write_tokens": -1`},
{name: "overflow", cacheWriteDetails: `"cached_tokens": 0, "cache_write_tokens": 9223372036854775808`},
}

for _, tt := range tests {
t.Run(tt.name, func(t *testing.T) {
resp := mustUnmarshalResponse(t, fmt.Sprintf(`{
"id": "resp_1",
"status": "completed",
"output": [],
"usage": {
"input_tokens": 100,
"input_tokens_details": {%s},
"output_tokens": 20,
"output_tokens_details": {"reasoning_tokens": 0},
"total_tokens": 120
}
}`, tt.cacheWriteDetails))

msg, err := toOutputMessage(resp, &model.Options{})
if err != nil {
t.Fatal(err)
}
if tokens, ok := GetCacheWriteTokens(msg); ok || tokens != 0 {
t.Fatalf("expected (0, false), got (%d, %v)", tokens, ok)
}
})
}
}

func TestCacheWriteTokensPreservedAfterConcat(t *testing.T) {
usageChunk := &schema.AgenticMessage{
Role: schema.AgenticRoleTypeAssistant,
Extra: map[string]any{keyOfCacheWriteTokens: 300},
}
textChunk := &schema.AgenticMessage{
Role: schema.AgenticRoleTypeAssistant,
ContentBlocks: []*schema.ContentBlock{
schema.NewContentBlockChunk(&schema.AssistantGenText{Text: "hello"}, &schema.StreamingMeta{Index: 0}),
},
}

msg, err := schema.ConcatAgenticMessages([]*schema.AgenticMessage{usageChunk, textChunk})
if err != nil {
t.Fatal(err)
}
if tokens, ok := GetCacheWriteTokens(msg); !ok || tokens != 300 {
t.Fatalf("expected (300, true), got (%d, %v)", tokens, ok)
}
}

func mustUnmarshalResponse(t *testing.T, raw string) *responses.Response {
t.Helper()
var resp responses.Response
if err := json.Unmarshal([]byte(raw), &resp); err != nil {
t.Fatalf("json.Unmarshal(response) error = %v", err)
}
return &resp
}
1 change: 1 addition & 0 deletions components/model/agenticopenai/responses_convertor.go
Original file line number Diff line number Diff line change
Expand Up @@ -1598,6 +1598,7 @@ func toOutputMessage(resp *responses.Response, options *model.Options) (msg *sch
Role: schema.AgenticRoleTypeAssistant,
ContentBlocks: blocks,
ResponseMeta: responseObjectToResponseMeta(resp),
Extra: cacheWriteTokensExtra(resp),
}

return msg, nil
Expand Down
28 changes: 14 additions & 14 deletions components/model/agenticopenai/responses_event_convertor.go
Original file line number Diff line number Diff line change
Expand Up @@ -48,24 +48,19 @@ func receivedStreamingResponse(sr *ssestream.Stream[responses.ResponseStreamEven
_ = sw.Send(nil, fmt.Errorf("received error event: code=%s message=%s", variant.Code, variant.Message))

case responses.ResponseCreatedEvent:
meta := responseObjectToResponseMeta(&variant.Response)
sender.sendMeta(meta, nil)
sender.sendResponse(&variant.Response, nil)

case responses.ResponseInProgressEvent:
meta := responseObjectToResponseMeta(&variant.Response)
sender.sendMeta(meta, nil)
sender.sendResponse(&variant.Response, nil)

case responses.ResponseCompletedEvent:
meta := responseObjectToResponseMeta(&variant.Response)
sender.sendMeta(meta, nil)
sender.sendResponse(&variant.Response, nil)

case responses.ResponseIncompleteEvent:
meta := responseObjectToResponseMeta(&variant.Response)
sender.sendMeta(meta, nil)
sender.sendResponse(&variant.Response, nil)

case responses.ResponseFailedEvent:
meta := responseObjectToResponseMeta(&variant.Response)
sender.sendMeta(meta, nil)
sender.sendResponse(&variant.Response, nil)

case responses.ResponseOutputItemAddedEvent:
blocks, err := receiver.itemAddedEventToContentBlock(variant)
Expand Down Expand Up @@ -224,17 +219,21 @@ func newCallbackSender(sw *schema.StreamWriter[*model.AgenticCallbackOutput], co
}
}

func (s *callbackSender) sendMeta(meta *schema.AgenticResponseMeta, err error) {
s.send(meta, nil, err)
func (s *callbackSender) sendResponse(resp *responses.Response, err error) {
if resp == nil {
s.send(nil, nil, nil, err)
return
}
s.send(responseObjectToResponseMeta(resp), nil, cacheWriteTokensExtra(resp), err)
}

func (s *callbackSender) sendBlock(block *schema.ContentBlock, err error) {
if block != nil || err != nil {
s.send(nil, block, err)
s.send(nil, block, nil, err)
}
}

func (s *callbackSender) send(meta *schema.AgenticResponseMeta, block *schema.ContentBlock, err error) {
func (s *callbackSender) send(meta *schema.AgenticResponseMeta, block *schema.ContentBlock, extra map[string]any, err error) {
if err != nil {
_ = s.sw.Send(nil, fmt.Errorf("%s: %w", s.errHeader, err))
return
Expand All @@ -243,6 +242,7 @@ func (s *callbackSender) send(meta *schema.AgenticResponseMeta, block *schema.Co
msg := &schema.AgenticMessage{
Role: schema.AgenticRoleTypeAssistant,
ResponseMeta: meta,
Extra: extra,
}

if block != nil {
Expand Down
16 changes: 1 addition & 15 deletions components/model/agenticopenai/responses_event_convertor_test.go
Original file line number Diff line number Diff line change
Expand Up @@ -764,20 +764,6 @@ func TestNewCallbackSender(t *testing.T) {
assert.Equal(t, config, s.config)
}

func TestCallbackSenderSendMeta(t *testing.T) {
sr, sw := schema.Pipe[*model.AgenticCallbackOutput](8)
r := sr.Copy(1)[0]
s := newCallbackSender(sw, &model.AgenticConfig{})

meta := &schema.AgenticResponseMeta{}
s.sendMeta(meta, nil)

out, err := r.Recv()
assert.NoError(t, err)
assert.NotNil(t, out)
assert.NotNil(t, out.Message.ResponseMeta)
}

func TestCallbackSenderSendBlock(t *testing.T) {
sr, sw := schema.Pipe[*model.AgenticCallbackOutput](8)
r := sr.Copy(1)[0]
Expand All @@ -798,7 +784,7 @@ func TestCallbackSenderSendError(t *testing.T) {
s := newCallbackSender(sw, &model.AgenticConfig{})
s.errHeader = "test error"

s.sendMeta(nil, errors.New("error"))
s.sendResponse(nil, errors.New("error"))

_, err := r.Recv()
assert.Error(t, err)
Expand Down
30 changes: 22 additions & 8 deletions components/model/agenticopenai/responses_model_test.go
Original file line number Diff line number Diff line change
Expand Up @@ -124,22 +124,36 @@ func TestModelStream(t *testing.T) {

mockey.Mock((*responses.ResponseService).NewStreaming).Return(mockStream).Build()

// Mock AsAny to return a completed event
completedResponse := mustUnmarshalResponse(t, `{
"id": "resp_1",
"status": "completed",
"output": [],
"usage": {
"input_tokens": 400,
"input_tokens_details": {"cached_tokens": 100, "cache_write_tokens": 300},
"output_tokens": 20,
"output_tokens_details": {"reasoning_tokens": 0},
"total_tokens": 420
}
}`)

// Mock AsAny to return a completed event with cache-write usage.
mockey.Mock(responses.ResponseStreamEventUnion.AsAny).Return(responses.ResponseCompletedEvent{
Response: responses.Response{
Output: []responses.ResponseOutputItemUnion{
{Type: "message", ID: "m1", Status: "completed"},
},
},
Response: *completedResponse,
}).Build()

s, err := m.Stream(ctx, input)
assert.NoError(t, err)
assert.NotNil(t, s)
defer s.Close()

// The stream should eventually close without errors
// We just verify it was created successfully
chunk, err := s.Recv()
assert.NoError(t, err)
if assert.NotNil(t, chunk) {
tokens, ok := GetCacheWriteTokens(chunk)
assert.True(t, ok)
assert.Equal(t, 300, tokens)
}
})

mockey.PatchConvey("genRequest error", func() {
Expand Down
Loading