101 lines
3.1 KiB
Go
101 lines
3.1 KiB
Go
package superagent
|
|
|
|
import (
|
|
"context"
|
|
"errors"
|
|
"fmt"
|
|
|
|
"fire-safety-ymd/internal/service"
|
|
)
|
|
|
|
// ChatAgentAdapter projects the SuperAgent-specific client into the
|
|
// provider-neutral service boundary.
|
|
type ChatAgentAdapter struct {
|
|
client Client
|
|
}
|
|
|
|
var _ service.ChatAgent = (*ChatAgentAdapter)(nil)
|
|
|
|
// NewChatAgentAdapter constructs the chat service adapter.
|
|
func NewChatAgentAdapter(client Client) (*ChatAgentAdapter, error) {
|
|
if client == nil {
|
|
return nil, errors.New("superagent client is required")
|
|
}
|
|
return &ChatAgentAdapter{client: client}, nil
|
|
}
|
|
|
|
// CreateSession maps only the server-controlled fields required by SuperAgent.
|
|
func (a *ChatAgentAdapter) CreateSession(ctx context.Context, request service.AgentCreateSessionRequest) (service.AgentSession, error) {
|
|
session, err := a.client.CreateSession(ctx, CreateSessionRequest{
|
|
ExternalSubjectID: request.ExternalSubjectID,
|
|
IdempotencyKey: request.IdempotencyKey,
|
|
RequestID: request.RequestID,
|
|
Metadata: request.Metadata,
|
|
})
|
|
if err != nil {
|
|
return service.AgentSession{}, mapChatError("create session", err)
|
|
}
|
|
return service.AgentSession{ID: session.ID}, nil
|
|
}
|
|
|
|
// StreamMessage removes provider identifiers, text deltas, and raw trace data
|
|
// from progress events. The final answer is returned only when the underlying
|
|
// strict SuperAgent client reports success.
|
|
func (a *ChatAgentAdapter) StreamMessage(
|
|
ctx context.Context,
|
|
request service.AgentMessageRequest,
|
|
traceHandler func(service.AgentTraceEvent),
|
|
) (service.AgentMessageResult, error) {
|
|
result, err := a.client.StreamMessage(ctx, StreamMessageRequest{
|
|
SessionID: request.SessionID,
|
|
Message: request.Message,
|
|
IdempotencyKey: request.IdempotencyKey,
|
|
RequestID: request.RequestID,
|
|
Metadata: request.Metadata,
|
|
}, func(event TraceEvent) {
|
|
if traceHandler == nil {
|
|
return
|
|
}
|
|
traceHandler(service.AgentTraceEvent{
|
|
Event: safeValue(event.Event),
|
|
ToolName: safeValue(event.ToolName),
|
|
Status: safeValue(event.Status),
|
|
})
|
|
})
|
|
if err != nil {
|
|
return service.AgentMessageResult{}, mapChatError("stream message", err)
|
|
}
|
|
return service.AgentMessageResult{
|
|
RunID: result.RunID,
|
|
ModelID: safeValue(result.ModelName),
|
|
Answer: result.Answer,
|
|
Usage: service.ChatTokenUsage{
|
|
Input: result.Usage.Input,
|
|
Output: result.Usage.Output,
|
|
Total: result.Usage.Total,
|
|
},
|
|
}, nil
|
|
}
|
|
|
|
func mapChatError(operation string, err error) error {
|
|
if err == nil {
|
|
return nil
|
|
}
|
|
if errors.Is(err, context.Canceled) || errors.Is(err, context.DeadlineExceeded) {
|
|
return err
|
|
}
|
|
var public error
|
|
switch {
|
|
case errors.Is(err, ErrRunFailed):
|
|
public = service.ErrChatRunFailed
|
|
case errors.Is(err, ErrProtocol), errors.Is(err, ErrStreamIncomplete), errors.Is(err, ErrStreamRead),
|
|
errors.Is(err, ErrRecoveryExhausted), errors.Is(err, ErrInvalidRequest):
|
|
public = service.ErrChatUpstreamProtocol
|
|
case errors.Is(err, ErrDisabled), errors.Is(err, ErrHTTPStatus):
|
|
public = service.ErrChatUpstreamUnavailable
|
|
default:
|
|
public = service.ErrChatUpstreamUnavailable
|
|
}
|
|
return fmt.Errorf("superagent %s: %w", operation, public)
|
|
}
|