iop/apps/node/internal/node/tunnel_handler.go

197 lines
5.2 KiB
Go

package node
import (
"context"
"fmt"
"time"
"go.uber.org/zap"
"iop/apps/node/internal/transport"
runtime "iop/packages/go/agentruntime"
iop "iop/proto/gen/iop"
)
// OnProviderTunnelRequest handles an incoming ProviderTunnelRequest from a transport Session.
func (n *Node) OnProviderTunnelRequest(ctx context.Context, sess *transport.Session, req *iop.ProviderTunnelRequest) error {
n.logger.Info("provider tunnel request received",
zap.String("run_id", req.GetRunId()),
zap.String("tunnel_id", req.GetTunnelId()),
zap.String("adapter", req.GetAdapter()),
zap.String("target", req.GetTarget()),
)
tr := runtime.ProviderTunnelRequest{
RunID: req.GetRunId(),
TunnelID: req.GetTunnelId(),
Adapter: req.GetAdapter(),
Target: req.GetTarget(),
Method: req.GetMethod(),
Path: req.GetPath(),
Headers: req.GetHeaders(),
Body: req.GetBody(),
Stream: req.GetStream(),
TimeoutSec: int(req.GetTimeoutSec()),
Metadata: req.GetMetadata(),
SessionID: req.GetSessionId(),
}
n.configSetMu.RLock()
configLocked := true
defer func() {
if configLocked {
n.configSetMu.RUnlock()
}
}()
adapter, err := n.router.LookupAdapter(tr.Adapter)
if err != nil {
n.sendTunnelError(sess, tr, fmt.Errorf("node: lookup adapter: %w", err))
return fmt.Errorf("node: lookup adapter: %w", err)
}
tunnelAdapter, ok := adapter.(runtime.ProviderTunnelAdapter)
if !ok {
err := fmt.Errorf("node: adapter %q does not support tunneling", tr.Adapter)
n.sendTunnelError(sess, tr, err)
return err
}
var caps runtime.Capabilities
if c, capsErr := adapter.Capabilities(ctx); capsErr == nil {
caps = c
}
admission := n.admissionFor(tr.Adapter, caps)
ticket, err := admission.acquire()
if err != nil {
n.logger.Warn("provider tunnel admission rejected",
zap.String("run_id", tr.RunID),
zap.String("tunnel_id", tr.TunnelID),
zap.String("adapter", tr.Adapter),
zap.String("reason", string(admissionRejectReason(err))),
)
n.sendTunnelError(sess, tr, err)
return fmt.Errorf("node: provider tunnel %s: %w", tr.TunnelID, err)
}
defer ticket.release()
var sender protoSender = noopSender{}
nodeID := n.nodeID
nodeAlias := ""
if sess != nil {
if sess.IsAlive() {
sender = sess
}
nodeID = sess.NodeID()
nodeAlias = sess.Alias()
}
sink := &tunnelSink{
sess: sender,
nodeID: nodeID,
nodeAlias: nodeAlias,
}
execCtx, cancel := context.WithCancel(ctx)
if tr.TimeoutSec > 0 {
execCtx, cancel = context.WithTimeout(ctx, time.Duration(tr.TimeoutSec)*time.Second)
}
defer cancel()
h := &runHandle{
runID: tr.RunID,
adapter: tr.Adapter,
target: tr.Target,
sessionID: normalizeSessionID(tr.SessionID),
cancel: cancel,
done: make(chan struct{}),
}
n.runs.register(h)
defer n.runs.deregister(tr.RunID)
defer close(h.done)
configLocked = false
n.configSetMu.RUnlock()
if err := tunnelAdapter.TunnelProvider(execCtx, tr, sink); err != nil {
n.logger.Warn("provider tunnel error",
zap.String("run_id", tr.RunID),
zap.String("tunnel_id", tr.TunnelID),
zap.Error(err),
)
return err
}
return nil
}
func (n *Node) sendTunnelError(sess *transport.Session, tr runtime.ProviderTunnelRequest, err error) {
if sess == nil || !sess.IsAlive() {
return
}
tf := &iop.ProviderTunnelFrame{
RunId: tr.RunID,
TunnelId: tr.TunnelID,
Sequence: 0,
Kind: iop.ProviderTunnelFrameKind_PROVIDER_TUNNEL_FRAME_KIND_ERROR,
Error: err.Error(),
Timestamp: time.Now().UnixNano(),
NodeId: sess.NodeID(),
NodeAlias: sess.Alias(),
}
_ = sess.Send(tf)
}
type tunnelSink struct {
sess protoSender
nodeID string
nodeAlias string
}
func (s *tunnelSink) EmitTunnelFrame(ctx context.Context, frame runtime.ProviderTunnelFrame) error {
var usage *iop.Usage
if frame.Usage != nil {
usage = &iop.Usage{
InputTokens: int32(frame.Usage.InputTokens),
OutputTokens: int32(frame.Usage.OutputTokens),
ReasoningTokens: int32(frame.Usage.ReasoningTokens),
CachedInputTokens: int32(frame.Usage.CachedInputTokens),
}
}
protoKind := iop.ProviderTunnelFrameKind_PROVIDER_TUNNEL_FRAME_KIND_UNSPECIFIED
switch frame.Kind {
case runtime.ProviderTunnelFrameKindResponseStart:
protoKind = iop.ProviderTunnelFrameKind_PROVIDER_TUNNEL_FRAME_KIND_RESPONSE_START
case runtime.ProviderTunnelFrameKindBody:
protoKind = iop.ProviderTunnelFrameKind_PROVIDER_TUNNEL_FRAME_KIND_BODY
case runtime.ProviderTunnelFrameKindEnd:
protoKind = iop.ProviderTunnelFrameKind_PROVIDER_TUNNEL_FRAME_KIND_END
case runtime.ProviderTunnelFrameKindError:
protoKind = iop.ProviderTunnelFrameKind_PROVIDER_TUNNEL_FRAME_KIND_ERROR
case runtime.ProviderTunnelFrameKindUsage:
protoKind = iop.ProviderTunnelFrameKind_PROVIDER_TUNNEL_FRAME_KIND_USAGE
}
tf := &iop.ProviderTunnelFrame{
RunId: frame.RunID,
TunnelId: frame.TunnelID,
Sequence: frame.Sequence,
Kind: protoKind,
StatusCode: int32(frame.StatusCode),
Headers: frame.Headers,
Body: frame.Body,
End: frame.End,
Error: frame.Error,
Usage: usage,
Metadata: frame.Metadata,
Timestamp: frame.Timestamp.UnixNano(),
NodeId: s.nodeID,
NodeAlias: s.nodeAlias,
}
if s.sess != nil {
return s.sess.Send(tf)
}
return nil
}