mirror of
https://github.com/QuantumNous/new-api.git
synced 2026-09-12 15:21:09 +00:00
fix(aws): cancel Bedrock requests on client disconnect (#6589)
* fix(aws): cancel Bedrock requests on client disconnect * fix(billing): log effective usage billing path
This commit is contained in:
@@ -40,11 +40,24 @@ func getAwsErrorStatusCode(err error) int {
|
|||||||
return http.StatusInternalServerError
|
return http.StatusInternalServerError
|
||||||
}
|
}
|
||||||
|
|
||||||
func newAwsInvokeContext() (context.Context, context.CancelFunc) {
|
func newAwsInvokeContext(parent context.Context) (context.Context, context.CancelFunc) {
|
||||||
if common.RelayTimeout <= 0 {
|
if common.RelayTimeout <= 0 {
|
||||||
return context.Background(), func() {}
|
return context.WithCancel(parent)
|
||||||
}
|
}
|
||||||
return context.WithTimeout(context.Background(), time.Duration(common.RelayTimeout)*time.Second)
|
return context.WithTimeout(parent, time.Duration(common.RelayTimeout)*time.Second)
|
||||||
|
}
|
||||||
|
|
||||||
|
func newAwsInvokeError(requestContext context.Context, err error, operation string) *types.NewAPIError {
|
||||||
|
options := make([]types.NewAPIErrorOptions, 0, 1)
|
||||||
|
if requestContext.Err() != nil {
|
||||||
|
options = append(options, types.ErrOptionWithSkipRetry())
|
||||||
|
}
|
||||||
|
return types.NewOpenAIError(
|
||||||
|
errors.Wrap(err, operation),
|
||||||
|
types.ErrorCodeAwsInvokeError,
|
||||||
|
getAwsErrorStatusCode(err),
|
||||||
|
options...,
|
||||||
|
)
|
||||||
}
|
}
|
||||||
|
|
||||||
func newAwsClient(c *gin.Context, info *relaycommon.RelayInfo) (*bedrockruntime.Client, error) {
|
func newAwsClient(c *gin.Context, info *relaycommon.RelayInfo) (*bedrockruntime.Client, error) {
|
||||||
@@ -215,13 +228,13 @@ func getAwsModelID(requestModel string) string {
|
|||||||
|
|
||||||
func awsHandler(c *gin.Context, info *relaycommon.RelayInfo, a *Adaptor) (*types.NewAPIError, *dto.Usage) {
|
func awsHandler(c *gin.Context, info *relaycommon.RelayInfo, a *Adaptor) (*types.NewAPIError, *dto.Usage) {
|
||||||
|
|
||||||
ctx, cancel := newAwsInvokeContext()
|
requestContext := c.Request.Context()
|
||||||
|
ctx, cancel := newAwsInvokeContext(requestContext)
|
||||||
defer cancel()
|
defer cancel()
|
||||||
|
|
||||||
awsResp, err := a.AwsClient.InvokeModel(ctx, a.AwsReq.(*bedrockruntime.InvokeModelInput))
|
awsResp, err := a.AwsClient.InvokeModel(ctx, a.AwsReq.(*bedrockruntime.InvokeModelInput))
|
||||||
if err != nil {
|
if err != nil {
|
||||||
statusCode := getAwsErrorStatusCode(err)
|
return newAwsInvokeError(requestContext, err, "InvokeModel"), nil
|
||||||
return types.NewOpenAIError(errors.Wrap(err, "InvokeModel"), types.ErrorCodeAwsInvokeError, statusCode), nil
|
|
||||||
}
|
}
|
||||||
|
|
||||||
claudeInfo := &claude.ClaudeResponseInfo{
|
claudeInfo := &claude.ClaudeResponseInfo{
|
||||||
@@ -245,13 +258,13 @@ func awsHandler(c *gin.Context, info *relaycommon.RelayInfo, a *Adaptor) (*types
|
|||||||
}
|
}
|
||||||
|
|
||||||
func awsStreamHandler(c *gin.Context, info *relaycommon.RelayInfo, a *Adaptor) (*types.NewAPIError, *dto.Usage) {
|
func awsStreamHandler(c *gin.Context, info *relaycommon.RelayInfo, a *Adaptor) (*types.NewAPIError, *dto.Usage) {
|
||||||
ctx, cancel := newAwsInvokeContext()
|
requestContext := c.Request.Context()
|
||||||
|
ctx, cancel := newAwsInvokeContext(requestContext)
|
||||||
defer cancel()
|
defer cancel()
|
||||||
|
|
||||||
awsResp, err := a.AwsClient.InvokeModelWithResponseStream(ctx, a.AwsReq.(*bedrockruntime.InvokeModelWithResponseStreamInput))
|
awsResp, err := a.AwsClient.InvokeModelWithResponseStream(ctx, a.AwsReq.(*bedrockruntime.InvokeModelWithResponseStreamInput))
|
||||||
if err != nil {
|
if err != nil {
|
||||||
statusCode := getAwsErrorStatusCode(err)
|
return newAwsInvokeError(requestContext, err, "InvokeModelWithResponseStream"), nil
|
||||||
return types.NewOpenAIError(errors.Wrap(err, "InvokeModelWithResponseStream"), types.ErrorCodeAwsInvokeError, statusCode), nil
|
|
||||||
}
|
}
|
||||||
stream := awsResp.GetStream()
|
stream := awsResp.GetStream()
|
||||||
defer stream.Close()
|
defer stream.Close()
|
||||||
@@ -264,23 +277,38 @@ func awsStreamHandler(c *gin.Context, info *relaycommon.RelayInfo, a *Adaptor) (
|
|||||||
Usage: &dto.Usage{},
|
Usage: &dto.Usage{},
|
||||||
}
|
}
|
||||||
|
|
||||||
for event := range stream.Events() {
|
events := stream.Events()
|
||||||
switch v := event.(type) {
|
streamLoop:
|
||||||
case *bedrockruntimeTypes.ResponseStreamMemberChunk:
|
for {
|
||||||
info.SetFirstResponseTime()
|
select {
|
||||||
respErr := claude.HandleStreamResponseData(c, info, claudeInfo, string(v.Value.Bytes))
|
case <-ctx.Done():
|
||||||
if respErr != nil {
|
break streamLoop
|
||||||
return respErr, nil
|
case event, ok := <-events:
|
||||||
|
if !ok {
|
||||||
|
break streamLoop
|
||||||
|
}
|
||||||
|
if ctx.Err() != nil {
|
||||||
|
break streamLoop
|
||||||
|
}
|
||||||
|
|
||||||
|
switch v := event.(type) {
|
||||||
|
case *bedrockruntimeTypes.ResponseStreamMemberChunk:
|
||||||
|
info.SetFirstResponseTime()
|
||||||
|
respErr := claude.HandleStreamResponseData(c, info, claudeInfo, string(v.Value.Bytes))
|
||||||
|
if respErr != nil {
|
||||||
|
return respErr, nil
|
||||||
|
}
|
||||||
|
case *bedrockruntimeTypes.UnknownUnionMember:
|
||||||
|
fmt.Println("unknown tag:", v.Tag)
|
||||||
|
return types.NewError(errors.New("unknown response type"), types.ErrorCodeInvalidRequest), nil
|
||||||
|
default:
|
||||||
|
fmt.Println("union is nil or unknown type")
|
||||||
|
return types.NewError(errors.New("nil or unknown response type"), types.ErrorCodeInvalidRequest), nil
|
||||||
}
|
}
|
||||||
case *bedrockruntimeTypes.UnknownUnionMember:
|
|
||||||
fmt.Println("unknown tag:", v.Tag)
|
|
||||||
return types.NewError(errors.New("unknown response type"), types.ErrorCodeInvalidRequest), nil
|
|
||||||
default:
|
|
||||||
fmt.Println("union is nil or unknown type")
|
|
||||||
return types.NewError(errors.New("nil or unknown response type"), types.ErrorCodeInvalidRequest), nil
|
|
||||||
}
|
}
|
||||||
}
|
}
|
||||||
|
|
||||||
|
_ = stream.Close()
|
||||||
claude.HandleStreamFinalResponse(c, info, claudeInfo)
|
claude.HandleStreamFinalResponse(c, info, claudeInfo)
|
||||||
return nil, claudeInfo.Usage
|
return nil, claudeInfo.Usage
|
||||||
}
|
}
|
||||||
@@ -288,13 +316,13 @@ func awsStreamHandler(c *gin.Context, info *relaycommon.RelayInfo, a *Adaptor) (
|
|||||||
// Nova模型处理函数
|
// Nova模型处理函数
|
||||||
func handleNovaRequest(c *gin.Context, info *relaycommon.RelayInfo, a *Adaptor) (*types.NewAPIError, *dto.Usage) {
|
func handleNovaRequest(c *gin.Context, info *relaycommon.RelayInfo, a *Adaptor) (*types.NewAPIError, *dto.Usage) {
|
||||||
|
|
||||||
ctx, cancel := newAwsInvokeContext()
|
requestContext := c.Request.Context()
|
||||||
|
ctx, cancel := newAwsInvokeContext(requestContext)
|
||||||
defer cancel()
|
defer cancel()
|
||||||
|
|
||||||
awsResp, err := a.AwsClient.InvokeModel(ctx, a.AwsReq.(*bedrockruntime.InvokeModelInput))
|
awsResp, err := a.AwsClient.InvokeModel(ctx, a.AwsReq.(*bedrockruntime.InvokeModelInput))
|
||||||
if err != nil {
|
if err != nil {
|
||||||
statusCode := getAwsErrorStatusCode(err)
|
return newAwsInvokeError(requestContext, err, "InvokeModel"), nil
|
||||||
return types.NewOpenAIError(errors.Wrap(err, "InvokeModel"), types.ErrorCodeAwsInvokeError, statusCode), nil
|
|
||||||
}
|
}
|
||||||
|
|
||||||
// 解析Nova响应
|
// 解析Nova响应
|
||||||
|
|||||||
@@ -2,17 +2,145 @@ package aws
|
|||||||
|
|
||||||
import (
|
import (
|
||||||
"bytes"
|
"bytes"
|
||||||
|
"context"
|
||||||
|
"errors"
|
||||||
|
"io"
|
||||||
"net/http"
|
"net/http"
|
||||||
"net/http/httptest"
|
"net/http/httptest"
|
||||||
|
"sync"
|
||||||
"testing"
|
"testing"
|
||||||
|
"time"
|
||||||
|
|
||||||
"github.com/QuantumNous/new-api/common"
|
"github.com/QuantumNous/new-api/common"
|
||||||
relaycommon "github.com/QuantumNous/new-api/relay/common"
|
relaycommon "github.com/QuantumNous/new-api/relay/common"
|
||||||
|
"github.com/QuantumNous/new-api/relaykit/dto"
|
||||||
|
relaytypes "github.com/QuantumNous/new-api/relaykit/types"
|
||||||
|
"github.com/aws/aws-sdk-go-v2/aws"
|
||||||
|
"github.com/aws/aws-sdk-go-v2/aws/protocol/eventstream"
|
||||||
|
"github.com/aws/aws-sdk-go-v2/aws/protocol/eventstream/eventstreamapi"
|
||||||
|
"github.com/aws/aws-sdk-go-v2/credentials"
|
||||||
"github.com/aws/aws-sdk-go-v2/service/bedrockruntime"
|
"github.com/aws/aws-sdk-go-v2/service/bedrockruntime"
|
||||||
"github.com/gin-gonic/gin"
|
"github.com/gin-gonic/gin"
|
||||||
|
"github.com/stretchr/testify/assert"
|
||||||
"github.com/stretchr/testify/require"
|
"github.com/stretchr/testify/require"
|
||||||
)
|
)
|
||||||
|
|
||||||
|
const awsTestModel = "anthropic.claude-3-5-sonnet-20240620-v1:0"
|
||||||
|
|
||||||
|
type awsHTTPClientFunc func(*http.Request) (*http.Response, error)
|
||||||
|
|
||||||
|
func (f awsHTTPClientFunc) Do(request *http.Request) (*http.Response, error) {
|
||||||
|
return f(request)
|
||||||
|
}
|
||||||
|
|
||||||
|
type awsNotifyingResponseWriter struct {
|
||||||
|
*httptest.ResponseRecorder
|
||||||
|
notifyOn []byte
|
||||||
|
notified chan int
|
||||||
|
once sync.Once
|
||||||
|
}
|
||||||
|
|
||||||
|
func newAwsNotifyingResponseWriter(notifyOn string) *awsNotifyingResponseWriter {
|
||||||
|
return &awsNotifyingResponseWriter{
|
||||||
|
ResponseRecorder: httptest.NewRecorder(),
|
||||||
|
notifyOn: []byte(notifyOn),
|
||||||
|
notified: make(chan int, 1),
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
func (w *awsNotifyingResponseWriter) Write(data []byte) (int, error) {
|
||||||
|
return w.ResponseRecorder.Write(data)
|
||||||
|
}
|
||||||
|
|
||||||
|
func (w *awsNotifyingResponseWriter) Flush() {
|
||||||
|
w.ResponseRecorder.Flush()
|
||||||
|
if bytes.Contains(w.Body.Bytes(), w.notifyOn) {
|
||||||
|
w.once.Do(func() {
|
||||||
|
w.notified <- w.Body.Len()
|
||||||
|
})
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
func newAwsTestClient(httpClient bedrockruntime.HTTPClient) *bedrockruntime.Client {
|
||||||
|
return bedrockruntime.New(bedrockruntime.Options{
|
||||||
|
Region: "us-east-1",
|
||||||
|
BaseEndpoint: aws.String("https://bedrock.test"),
|
||||||
|
Credentials: aws.NewCredentialsCache(credentials.NewStaticCredentialsProvider(
|
||||||
|
"access-key", "secret-key", "",
|
||||||
|
)),
|
||||||
|
HTTPClient: httpClient,
|
||||||
|
Retryer: aws.NopRetryer{},
|
||||||
|
})
|
||||||
|
}
|
||||||
|
|
||||||
|
func newAwsTestContext(writer http.ResponseWriter, requestContext context.Context) *gin.Context {
|
||||||
|
c, _ := gin.CreateTestContext(writer)
|
||||||
|
c.Request = httptest.NewRequest(http.MethodPost, "/v1/chat/completions", nil).WithContext(requestContext)
|
||||||
|
return c
|
||||||
|
}
|
||||||
|
|
||||||
|
func newAwsTestRelayInfo() *relaycommon.RelayInfo {
|
||||||
|
return &relaycommon.RelayInfo{
|
||||||
|
StartTime: time.Now(),
|
||||||
|
IsStream: true,
|
||||||
|
OriginModelName: awsTestModel,
|
||||||
|
RelayFormat: relaytypes.RelayFormatOpenAI,
|
||||||
|
ShouldIncludeUsage: true,
|
||||||
|
ChannelMeta: &relaycommon.ChannelMeta{
|
||||||
|
UpstreamModelName: awsTestModel,
|
||||||
|
},
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
func newAwsInvokeModelInput() *bedrockruntime.InvokeModelInput {
|
||||||
|
return &bedrockruntime.InvokeModelInput{
|
||||||
|
ModelId: aws.String(awsTestModel),
|
||||||
|
Body: []byte(`{}`),
|
||||||
|
Accept: aws.String("application/json"),
|
||||||
|
ContentType: aws.String("application/json"),
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
func newAwsStreamInput() *bedrockruntime.InvokeModelWithResponseStreamInput {
|
||||||
|
return &bedrockruntime.InvokeModelWithResponseStreamInput{
|
||||||
|
ModelId: aws.String(awsTestModel),
|
||||||
|
Body: []byte(`{}`),
|
||||||
|
Accept: aws.String("application/json"),
|
||||||
|
ContentType: aws.String("application/json"),
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
func writeAwsStreamEvent(writer io.Writer, data string) error {
|
||||||
|
payload, err := common.Marshal(struct {
|
||||||
|
Bytes []byte `json:"bytes"`
|
||||||
|
}{Bytes: []byte(data)})
|
||||||
|
if err != nil {
|
||||||
|
return err
|
||||||
|
}
|
||||||
|
|
||||||
|
return eventstream.NewEncoder().Encode(writer, eventstream.Message{
|
||||||
|
Headers: eventstream.Headers{
|
||||||
|
{Name: eventstreamapi.MessageTypeHeader, Value: eventstream.StringValue(eventstreamapi.EventMessageType)},
|
||||||
|
{Name: eventstreamapi.EventTypeHeader, Value: eventstream.StringValue("chunk")},
|
||||||
|
{Name: eventstreamapi.ContentTypeHeader, Value: eventstream.StringValue("application/json")},
|
||||||
|
},
|
||||||
|
Payload: payload,
|
||||||
|
})
|
||||||
|
}
|
||||||
|
|
||||||
|
func newAwsStreamResponse(request *http.Request, body io.ReadCloser) *http.Response {
|
||||||
|
return &http.Response{
|
||||||
|
StatusCode: http.StatusOK,
|
||||||
|
Status: "200 OK",
|
||||||
|
Header: http.Header{
|
||||||
|
"Content-Type": []string{"application/vnd.amazon.eventstream"},
|
||||||
|
"X-Amzn-Bedrock-Content-Type": []string{"application/json"},
|
||||||
|
},
|
||||||
|
Body: body,
|
||||||
|
Request: request,
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
func TestDoAwsClientRequest_AppliesRuntimeHeaderOverrideToAnthropicBeta(t *testing.T) {
|
func TestDoAwsClientRequest_AppliesRuntimeHeaderOverrideToAnthropicBeta(t *testing.T) {
|
||||||
t.Parallel()
|
t.Parallel()
|
||||||
|
|
||||||
@@ -53,3 +181,278 @@ func TestDoAwsClientRequest_AppliesRuntimeHeaderOverrideToAnthropicBeta(t *testi
|
|||||||
require.True(t, ok)
|
require.True(t, ok)
|
||||||
require.Equal(t, []any{"computer-use-2025-01-24"}, values)
|
require.Equal(t, []any{"computer-use-2025-01-24"}, values)
|
||||||
}
|
}
|
||||||
|
|
||||||
|
func TestNewAwsInvokeContextInheritsParent(t *testing.T) {
|
||||||
|
originalRelayTimeout := common.RelayTimeout
|
||||||
|
t.Cleanup(func() {
|
||||||
|
common.RelayTimeout = originalRelayTimeout
|
||||||
|
})
|
||||||
|
|
||||||
|
tests := []struct {
|
||||||
|
name string
|
||||||
|
relayTimeout int
|
||||||
|
wantDeadline bool
|
||||||
|
}{
|
||||||
|
{name: "without relay timeout", relayTimeout: 0, wantDeadline: false},
|
||||||
|
{name: "with relay timeout", relayTimeout: 30, wantDeadline: true},
|
||||||
|
}
|
||||||
|
|
||||||
|
for _, test := range tests {
|
||||||
|
t.Run(test.name, func(t *testing.T) {
|
||||||
|
common.RelayTimeout = test.relayTimeout
|
||||||
|
parent, cancelParent := context.WithCancel(context.Background())
|
||||||
|
invokeContext, cancelInvoke := newAwsInvokeContext(parent)
|
||||||
|
defer cancelInvoke()
|
||||||
|
|
||||||
|
_, hasDeadline := invokeContext.Deadline()
|
||||||
|
assert.Equal(t, test.wantDeadline, hasDeadline)
|
||||||
|
|
||||||
|
cancelParent()
|
||||||
|
require.ErrorIs(t, invokeContext.Err(), context.Canceled)
|
||||||
|
})
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
func TestNewAwsInvokeErrorSkipsRetryOnlyForClientCancellation(t *testing.T) {
|
||||||
|
canceledContext, cancel := context.WithCancel(context.Background())
|
||||||
|
cancel()
|
||||||
|
|
||||||
|
tests := []struct {
|
||||||
|
name string
|
||||||
|
requestContext context.Context
|
||||||
|
err error
|
||||||
|
wantSkipRetry bool
|
||||||
|
}{
|
||||||
|
{
|
||||||
|
name: "client context canceled",
|
||||||
|
requestContext: canceledContext,
|
||||||
|
err: context.Canceled,
|
||||||
|
wantSkipRetry: true,
|
||||||
|
},
|
||||||
|
{
|
||||||
|
name: "relay timeout with live client context",
|
||||||
|
requestContext: context.Background(),
|
||||||
|
err: context.DeadlineExceeded,
|
||||||
|
wantSkipRetry: false,
|
||||||
|
},
|
||||||
|
{
|
||||||
|
name: "upstream error with live client context",
|
||||||
|
requestContext: context.Background(),
|
||||||
|
err: errors.New("upstream failed"),
|
||||||
|
wantSkipRetry: false,
|
||||||
|
},
|
||||||
|
}
|
||||||
|
|
||||||
|
for _, test := range tests {
|
||||||
|
t.Run(test.name, func(t *testing.T) {
|
||||||
|
err := newAwsInvokeError(test.requestContext, test.err, "InvokeModel")
|
||||||
|
assert.Equal(t, test.wantSkipRetry, relaytypes.IsSkipRetryError(err))
|
||||||
|
})
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
func TestAwsHandlersCancelSdkRequestAndSkipRetry(t *testing.T) {
|
||||||
|
originalRelayTimeout := common.RelayTimeout
|
||||||
|
common.RelayTimeout = 0
|
||||||
|
t.Cleanup(func() {
|
||||||
|
common.RelayTimeout = originalRelayTimeout
|
||||||
|
})
|
||||||
|
|
||||||
|
tests := []struct {
|
||||||
|
name string
|
||||||
|
request any
|
||||||
|
handle func(*gin.Context, *relaycommon.RelayInfo, *Adaptor) (*relaytypes.NewAPIError, *dto.Usage)
|
||||||
|
}{
|
||||||
|
{name: "non-stream", request: newAwsInvokeModelInput(), handle: awsHandler},
|
||||||
|
{name: "stream", request: newAwsStreamInput(), handle: awsStreamHandler},
|
||||||
|
{name: "nova", request: newAwsInvokeModelInput(), handle: handleNovaRequest},
|
||||||
|
}
|
||||||
|
|
||||||
|
for _, test := range tests {
|
||||||
|
t.Run(test.name, func(t *testing.T) {
|
||||||
|
requestContext, cancelRequest := context.WithCancel(context.Background())
|
||||||
|
t.Cleanup(cancelRequest)
|
||||||
|
|
||||||
|
upstreamContexts := make(chan context.Context, 1)
|
||||||
|
client := newAwsTestClient(awsHTTPClientFunc(func(request *http.Request) (*http.Response, error) {
|
||||||
|
upstreamContexts <- request.Context()
|
||||||
|
<-request.Context().Done()
|
||||||
|
return nil, request.Context().Err()
|
||||||
|
}))
|
||||||
|
adaptor := &Adaptor{AwsClient: client, AwsReq: test.request}
|
||||||
|
c := newAwsTestContext(httptest.NewRecorder(), requestContext)
|
||||||
|
info := newAwsTestRelayInfo()
|
||||||
|
|
||||||
|
type handlerResult struct {
|
||||||
|
err *relaytypes.NewAPIError
|
||||||
|
usage *dto.Usage
|
||||||
|
}
|
||||||
|
results := make(chan handlerResult, 1)
|
||||||
|
go func() {
|
||||||
|
err, usage := test.handle(c, info, adaptor)
|
||||||
|
results <- handlerResult{err: err, usage: usage}
|
||||||
|
}()
|
||||||
|
|
||||||
|
var upstreamContext context.Context
|
||||||
|
select {
|
||||||
|
case upstreamContext = <-upstreamContexts:
|
||||||
|
case result := <-results:
|
||||||
|
t.Fatalf("handler returned before issuing AWS request: %v", result.err)
|
||||||
|
case <-time.After(5 * time.Second):
|
||||||
|
t.Fatal("AWS request did not start")
|
||||||
|
}
|
||||||
|
|
||||||
|
cancelRequest()
|
||||||
|
|
||||||
|
var result handlerResult
|
||||||
|
select {
|
||||||
|
case result = <-results:
|
||||||
|
case <-time.After(5 * time.Second):
|
||||||
|
t.Fatal("handler did not stop after client cancellation")
|
||||||
|
}
|
||||||
|
|
||||||
|
require.ErrorIs(t, upstreamContext.Err(), context.Canceled)
|
||||||
|
require.NotNil(t, result.err)
|
||||||
|
assert.True(t, relaytypes.IsSkipRetryError(result.err))
|
||||||
|
assert.Nil(t, result.usage)
|
||||||
|
})
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
func TestAwsStreamHandlerUsesFinalUpstreamUsage(t *testing.T) {
|
||||||
|
originalRelayTimeout := common.RelayTimeout
|
||||||
|
common.RelayTimeout = 0
|
||||||
|
t.Cleanup(func() {
|
||||||
|
common.RelayTimeout = originalRelayTimeout
|
||||||
|
})
|
||||||
|
|
||||||
|
events := []string{
|
||||||
|
`{"type":"message_start","message":{"id":"msg_test","type":"message","role":"assistant","model":"claude-test","content":[],"usage":{"input_tokens":100,"output_tokens":1}}}`,
|
||||||
|
`{"type":"content_block_start","index":0,"content_block":{"type":"text","text":""}}`,
|
||||||
|
`{"type":"content_block_delta","index":0,"delta":{"type":"text_delta","text":"partial"}}`,
|
||||||
|
`{"type":"message_delta","delta":{"stop_reason":"end_turn"},"usage":{"output_tokens":423}}`,
|
||||||
|
`{"type":"message_stop"}`,
|
||||||
|
}
|
||||||
|
client := newAwsTestClient(awsHTTPClientFunc(func(request *http.Request) (*http.Response, error) {
|
||||||
|
var body bytes.Buffer
|
||||||
|
for _, event := range events {
|
||||||
|
if err := writeAwsStreamEvent(&body, event); err != nil {
|
||||||
|
return nil, err
|
||||||
|
}
|
||||||
|
}
|
||||||
|
return newAwsStreamResponse(request, io.NopCloser(bytes.NewReader(body.Bytes()))), nil
|
||||||
|
}))
|
||||||
|
adaptor := &Adaptor{AwsClient: client, AwsReq: newAwsStreamInput()}
|
||||||
|
recorder := httptest.NewRecorder()
|
||||||
|
c := newAwsTestContext(recorder, context.Background())
|
||||||
|
|
||||||
|
handlerErr, usage := awsStreamHandler(c, newAwsTestRelayInfo(), adaptor)
|
||||||
|
|
||||||
|
require.Nil(t, handlerErr)
|
||||||
|
require.NotNil(t, usage)
|
||||||
|
require.NotNil(t, usage.BillingUsage)
|
||||||
|
require.NotNil(t, usage.BillingUsage.ClaudeUsage)
|
||||||
|
assert.Equal(t, 100, usage.BillingUsage.ClaudeUsage.InputTokens)
|
||||||
|
assert.Equal(t, 423, usage.BillingUsage.ClaudeUsage.OutputTokens)
|
||||||
|
assert.Contains(t, recorder.Body.String(), "[DONE]")
|
||||||
|
}
|
||||||
|
|
||||||
|
func TestAwsStreamHandlerStopsAtClientCancellationAndKeepsPartialBillingUsage(t *testing.T) {
|
||||||
|
originalRelayTimeout := common.RelayTimeout
|
||||||
|
common.RelayTimeout = 0
|
||||||
|
t.Cleanup(func() {
|
||||||
|
common.RelayTimeout = originalRelayTimeout
|
||||||
|
})
|
||||||
|
|
||||||
|
requestContext, cancelRequest := context.WithCancel(context.Background())
|
||||||
|
t.Cleanup(cancelRequest)
|
||||||
|
releaseFinal := make(chan struct{})
|
||||||
|
var releaseFinalOnce sync.Once
|
||||||
|
release := func() {
|
||||||
|
releaseFinalOnce.Do(func() {
|
||||||
|
close(releaseFinal)
|
||||||
|
})
|
||||||
|
}
|
||||||
|
t.Cleanup(release)
|
||||||
|
|
||||||
|
producerResults := make(chan error, 1)
|
||||||
|
upstreamContexts := make(chan context.Context, 1)
|
||||||
|
client := newAwsTestClient(awsHTTPClientFunc(func(request *http.Request) (*http.Response, error) {
|
||||||
|
upstreamContexts <- request.Context()
|
||||||
|
reader, writer := io.Pipe()
|
||||||
|
go func() {
|
||||||
|
defer writer.Close()
|
||||||
|
initialEvents := []string{
|
||||||
|
`{"type":"message_start","message":{"id":"msg_test","type":"message","role":"assistant","model":"claude-test","content":[],"usage":{"input_tokens":100,"output_tokens":1}}}`,
|
||||||
|
`{"type":"content_block_start","index":0,"content_block":{"type":"text","text":""}}`,
|
||||||
|
`{"type":"content_block_delta","index":0,"delta":{"type":"text_delta","text":"partial"}}`,
|
||||||
|
}
|
||||||
|
for _, event := range initialEvents {
|
||||||
|
if err := writeAwsStreamEvent(writer, event); err != nil {
|
||||||
|
producerResults <- err
|
||||||
|
return
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
<-releaseFinal
|
||||||
|
producerResults <- writeAwsStreamEvent(writer, `{"type":"message_delta","delta":{"stop_reason":"end_turn"},"usage":{"output_tokens":423}}`)
|
||||||
|
}()
|
||||||
|
return newAwsStreamResponse(request, reader), nil
|
||||||
|
}))
|
||||||
|
|
||||||
|
responseWriter := newAwsNotifyingResponseWriter("partial")
|
||||||
|
c := newAwsTestContext(responseWriter, requestContext)
|
||||||
|
adaptor := &Adaptor{AwsClient: client, AwsReq: newAwsStreamInput()}
|
||||||
|
|
||||||
|
type handlerResult struct {
|
||||||
|
err *relaytypes.NewAPIError
|
||||||
|
usage *dto.Usage
|
||||||
|
}
|
||||||
|
results := make(chan handlerResult, 1)
|
||||||
|
go func() {
|
||||||
|
err, usage := awsStreamHandler(c, newAwsTestRelayInfo(), adaptor)
|
||||||
|
results <- handlerResult{err: err, usage: usage}
|
||||||
|
}()
|
||||||
|
|
||||||
|
var upstreamContext context.Context
|
||||||
|
select {
|
||||||
|
case upstreamContext = <-upstreamContexts:
|
||||||
|
case <-time.After(5 * time.Second):
|
||||||
|
t.Fatal("AWS stream request did not start")
|
||||||
|
}
|
||||||
|
|
||||||
|
var bodyLengthBeforeCancel int
|
||||||
|
select {
|
||||||
|
case bodyLengthBeforeCancel = <-responseWriter.notified:
|
||||||
|
case <-time.After(5 * time.Second):
|
||||||
|
t.Fatal("partial response was not written")
|
||||||
|
}
|
||||||
|
cancelRequest()
|
||||||
|
|
||||||
|
var result handlerResult
|
||||||
|
select {
|
||||||
|
case result = <-results:
|
||||||
|
case <-time.After(5 * time.Second):
|
||||||
|
t.Fatal("stream handler did not stop after client cancellation")
|
||||||
|
}
|
||||||
|
|
||||||
|
require.ErrorIs(t, upstreamContext.Err(), context.Canceled)
|
||||||
|
require.Nil(t, result.err)
|
||||||
|
require.NotNil(t, result.usage)
|
||||||
|
require.NotNil(t, result.usage.BillingUsage)
|
||||||
|
require.NotNil(t, result.usage.BillingUsage.ClaudeUsage)
|
||||||
|
assert.Equal(t, dto.BillingUsageSourceClaudeMessages, result.usage.BillingUsage.Source)
|
||||||
|
assert.Equal(t, dto.BillingUsageSemanticAnthropic, result.usage.BillingUsage.Semantic)
|
||||||
|
assert.Equal(t, 100, result.usage.BillingUsage.ClaudeUsage.InputTokens)
|
||||||
|
assert.Equal(t, 1, result.usage.BillingUsage.ClaudeUsage.OutputTokens)
|
||||||
|
assert.Equal(t, bodyLengthBeforeCancel, responseWriter.Body.Len())
|
||||||
|
assert.NotContains(t, responseWriter.Body.String(), "[DONE]")
|
||||||
|
|
||||||
|
release()
|
||||||
|
select {
|
||||||
|
case producerErr := <-producerResults:
|
||||||
|
require.Error(t, producerErr)
|
||||||
|
case <-time.After(5 * time.Second):
|
||||||
|
t.Fatal("upstream producer did not observe the closed stream")
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|||||||
+11
-15
@@ -25,36 +25,32 @@ func effectiveBillingUsage(usage *dto.Usage) *dto.Usage {
|
|||||||
}
|
}
|
||||||
|
|
||||||
func usageBillingPathForLog(isLocalCountTokens bool, usage *dto.Usage) string {
|
func usageBillingPathForLog(isLocalCountTokens bool, usage *dto.Usage) string {
|
||||||
if isLocalCountTokens {
|
effectiveUsage, ok := usageFromBillingUsage(usage)
|
||||||
return usageBillingPathLocal
|
if !ok {
|
||||||
}
|
if isLocalCountTokens {
|
||||||
if usage == nil || usage.BillingUsage == nil {
|
return usageBillingPathLocal
|
||||||
|
}
|
||||||
return usageBillingPathUpstream
|
return usageBillingPathUpstream
|
||||||
}
|
}
|
||||||
source := strings.TrimSpace(usage.BillingUsage.Source)
|
|
||||||
semantic := strings.TrimSpace(usage.BillingUsage.Semantic)
|
switch effectiveUsage.UsageSemantic {
|
||||||
if strings.EqualFold(source, dto.BillingUsageSourceOAIChat) ||
|
case dto.BillingUsageSemanticOpenAI:
|
||||||
strings.EqualFold(source, dto.BillingUsageSourceOAIResponses) ||
|
|
||||||
strings.EqualFold(semantic, dto.BillingUsageSemanticOpenAI) {
|
|
||||||
if usage.BillingUsage.Estimated {
|
if usage.BillingUsage.Estimated {
|
||||||
return usageBillingPathOpenAIEstimated
|
return usageBillingPathOpenAIEstimated
|
||||||
}
|
}
|
||||||
return usageBillingPathOpenAI
|
return usageBillingPathOpenAI
|
||||||
}
|
case dto.BillingUsageSemanticAnthropic:
|
||||||
if strings.EqualFold(source, dto.BillingUsageSourceClaudeMessages) ||
|
|
||||||
strings.EqualFold(semantic, dto.BillingUsageSemanticAnthropic) {
|
|
||||||
if usage.BillingUsage.Estimated {
|
if usage.BillingUsage.Estimated {
|
||||||
return usageBillingPathAnthropicEstimated
|
return usageBillingPathAnthropicEstimated
|
||||||
}
|
}
|
||||||
return usageBillingPathAnthropic
|
return usageBillingPathAnthropic
|
||||||
}
|
case dto.BillingUsageSemanticGemini:
|
||||||
if strings.EqualFold(source, dto.BillingUsageSourceGeminiChat) ||
|
|
||||||
strings.EqualFold(semantic, dto.BillingUsageSemanticGemini) {
|
|
||||||
if usage.BillingUsage.Estimated {
|
if usage.BillingUsage.Estimated {
|
||||||
return usageBillingPathGeminiEstimated
|
return usageBillingPathGeminiEstimated
|
||||||
}
|
}
|
||||||
return usageBillingPathGemini
|
return usageBillingPathGemini
|
||||||
}
|
}
|
||||||
|
|
||||||
return usageBillingPathUpstream
|
return usageBillingPathUpstream
|
||||||
}
|
}
|
||||||
|
|
||||||
|
|||||||
@@ -284,9 +284,18 @@ func TestCalculateTextQuotaSummaryUsesOpenAIBillingUsageBeforeTopLevelUsage(t *t
|
|||||||
}
|
}
|
||||||
|
|
||||||
func TestUsageBillingPathForLog(t *testing.T) {
|
func TestUsageBillingPathForLog(t *testing.T) {
|
||||||
require.Equal(t, usageBillingPathLocal, usageBillingPathForLog(true, &dto.Usage{
|
require.Equal(t, usageBillingPathAnthropic, usageBillingPathForLog(true, &dto.Usage{
|
||||||
BillingUsage: dto.NewClaudeMessagesBillingUsage(&dto.ClaudeUsage{InputTokens: 1}),
|
BillingUsage: dto.NewClaudeMessagesBillingUsage(&dto.ClaudeUsage{InputTokens: 1}),
|
||||||
}))
|
}))
|
||||||
|
invalidBillingUsage := &dto.Usage{
|
||||||
|
PromptTokens: 1,
|
||||||
|
BillingUsage: &dto.BillingUsage{
|
||||||
|
Source: dto.BillingUsageSourceClaudeMessages,
|
||||||
|
Semantic: dto.BillingUsageSemanticAnthropic,
|
||||||
|
},
|
||||||
|
}
|
||||||
|
require.Equal(t, usageBillingPathLocal, usageBillingPathForLog(true, invalidBillingUsage))
|
||||||
|
require.Equal(t, usageBillingPathUpstream, usageBillingPathForLog(false, invalidBillingUsage))
|
||||||
require.Equal(t, usageBillingPathUpstream, usageBillingPathForLog(false, &dto.Usage{}))
|
require.Equal(t, usageBillingPathUpstream, usageBillingPathForLog(false, &dto.Usage{}))
|
||||||
require.Equal(t, usageBillingPathOpenAI, usageBillingPathForLog(false, &dto.Usage{
|
require.Equal(t, usageBillingPathOpenAI, usageBillingPathForLog(false, &dto.Usage{
|
||||||
BillingUsage: dto.NewOpenAIChatBillingUsage(&dto.Usage{PromptTokens: 1}),
|
BillingUsage: dto.NewOpenAIChatBillingUsage(&dto.Usage{PromptTokens: 1}),
|
||||||
@@ -297,7 +306,7 @@ func TestUsageBillingPathForLog(t *testing.T) {
|
|||||||
require.Equal(t, usageBillingPathGemini, usageBillingPathForLog(false, &dto.Usage{
|
require.Equal(t, usageBillingPathGemini, usageBillingPathForLog(false, &dto.Usage{
|
||||||
BillingUsage: dto.NewGeminiChatBillingUsage(&dto.GeminiUsageMetadata{PromptTokenCount: 1}),
|
BillingUsage: dto.NewGeminiChatBillingUsage(&dto.GeminiUsageMetadata{PromptTokenCount: 1}),
|
||||||
}))
|
}))
|
||||||
require.Equal(t, usageBillingPathGeminiEstimated, usageBillingPathForLog(false, &dto.Usage{
|
require.Equal(t, usageBillingPathGeminiEstimated, usageBillingPathForLog(true, &dto.Usage{
|
||||||
BillingUsage: dto.NewEstimatedGeminiChatBillingUsage(&dto.Usage{PromptTokens: 1}),
|
BillingUsage: dto.NewEstimatedGeminiChatBillingUsage(&dto.Usage{PromptTokens: 1}),
|
||||||
}))
|
}))
|
||||||
}
|
}
|
||||||
@@ -306,7 +315,7 @@ func TestAppendUsageBillingPathForLogWritesAdminInfo(t *testing.T) {
|
|||||||
other := map[string]interface{}{
|
other := map[string]interface{}{
|
||||||
"admin_info": map[string]interface{}{},
|
"admin_info": map[string]interface{}{},
|
||||||
}
|
}
|
||||||
appendUsageBillingPathForLog(other, false, &dto.Usage{
|
appendUsageBillingPathForLog(other, true, &dto.Usage{
|
||||||
BillingUsage: dto.NewClaudeMessagesBillingUsage(&dto.ClaudeUsage{InputTokens: 1}),
|
BillingUsage: dto.NewClaudeMessagesBillingUsage(&dto.ClaudeUsage{InputTokens: 1}),
|
||||||
})
|
})
|
||||||
|
|
||||||
|
|||||||
Reference in New Issue
Block a user