197 lines
5.2 KiB
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
|
|
}
|