Skip to content
Merged
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
78 changes: 53 additions & 25 deletions relay/channel/aws/relay-aws.go
Original file line number Diff line number Diff line change
Expand Up @@ -40,11 +40,24 @@ func getAwsErrorStatusCode(err error) int {
return http.StatusInternalServerError
}

func newAwsInvokeContext() (context.Context, context.CancelFunc) {
func newAwsInvokeContext(parent context.Context) (context.Context, context.CancelFunc) {
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) {
Expand Down Expand Up @@ -215,13 +228,13 @@ func getAwsModelID(requestModel string) string {

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()

awsResp, err := a.AwsClient.InvokeModel(ctx, a.AwsReq.(*bedrockruntime.InvokeModelInput))
if err != nil {
statusCode := getAwsErrorStatusCode(err)
return types.NewOpenAIError(errors.Wrap(err, "InvokeModel"), types.ErrorCodeAwsInvokeError, statusCode), nil
return newAwsInvokeError(requestContext, err, "InvokeModel"), nil
}

claudeInfo := &claude.ClaudeResponseInfo{
Expand All @@ -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) {
ctx, cancel := newAwsInvokeContext()
requestContext := c.Request.Context()
ctx, cancel := newAwsInvokeContext(requestContext)
defer cancel()

awsResp, err := a.AwsClient.InvokeModelWithResponseStream(ctx, a.AwsReq.(*bedrockruntime.InvokeModelWithResponseStreamInput))
if err != nil {
statusCode := getAwsErrorStatusCode(err)
return types.NewOpenAIError(errors.Wrap(err, "InvokeModelWithResponseStream"), types.ErrorCodeAwsInvokeError, statusCode), nil
return newAwsInvokeError(requestContext, err, "InvokeModelWithResponseStream"), nil
}
stream := awsResp.GetStream()
defer stream.Close()
Expand All @@ -264,37 +277,52 @@ func awsStreamHandler(c *gin.Context, info *relaycommon.RelayInfo, a *Adaptor) (
Usage: &dto.Usage{},
}

for event := range stream.Events() {
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
events := stream.Events()
streamLoop:
for {
select {
case <-ctx.Done():
break streamLoop
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)
return nil, claudeInfo.Usage
}

// Nova模型处理函数
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()

awsResp, err := a.AwsClient.InvokeModel(ctx, a.AwsReq.(*bedrockruntime.InvokeModelInput))
if err != nil {
statusCode := getAwsErrorStatusCode(err)
return types.NewOpenAIError(errors.Wrap(err, "InvokeModel"), types.ErrorCodeAwsInvokeError, statusCode), nil
return newAwsInvokeError(requestContext, err, "InvokeModel"), nil
}

// 解析Nova响应
Expand Down
Loading
Loading