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
5 changes: 5 additions & 0 deletions control-plane/internal/cli/session.go
Original file line number Diff line number Diff line change
Expand Up @@ -31,6 +31,7 @@ type sessionToolOptions struct {
}

type sessionOfferOptions struct {
target string
provider string
transport string
sdpSource string
Expand Down Expand Up @@ -113,6 +114,7 @@ func newSessionOfferCommand() *cobra.Command {
}
cmd.Flags().StringVar(&opts.provider, "provider", "", "Explicit session provider")
cmd.Flags().StringVar(&opts.transport, "transport", "", "Explicit session transport")
cmd.Flags().StringVar(&opts.target, "target", "", "Registered <node>.<session> whose turn detection settings to use")
cmd.Flags().StringVar(&opts.sdpSource, "sdp", "", "SDP offer as inline text, @path, or - for stdin; defaults to stdin")
cmd.Flags().StringVarP(&opts.outputFormat, "output", "o", "raw", "Output format: raw, json, pretty, yaml")
return cmd
Expand All @@ -133,6 +135,9 @@ func runSessionOffer(ctx context.Context, sessionID string, opts *sessionOfferOp
if strings.TrimSpace(opts.transport) != "" {
values.Set("transport", opts.transport)
}
if strings.TrimSpace(opts.target) != "" {
values.Set("target", opts.target)
}
path := "/api/v1/session-instances/" + url.PathEscape(sessionID) + "/realtime-offer"
if encoded := values.Encode(); encoded != "" {
path += "?" + encoded
Expand Down
2 changes: 2 additions & 0 deletions control-plane/internal/cli/session_test.go
Original file line number Diff line number Diff line change
Expand Up @@ -77,6 +77,7 @@ func TestRunSessionOfferPostsSDPAndWritesRawAnswer(t *testing.T) {
require.Equal(t, "/api/v1/session-instances/sess-1/realtime-offer", r.URL.Path)
require.Equal(t, "openai", r.URL.Query().Get("provider"))
require.Equal(t, "webrtc", r.URL.Query().Get("transport"))
require.Equal(t, "support.voice", r.URL.Query().Get("target"))
gotContentType = r.Header.Get("Content-Type")
gotAPIKey = r.Header.Get("X-API-Key")
body, err := io.ReadAll(r.Body)
Expand All @@ -90,6 +91,7 @@ func TestRunSessionOfferPostsSDPAndWritesRawAnswer(t *testing.T) {
var stdout bytes.Buffer
err := runSessionOffer(context.Background(), "sess-1", &sessionOfferOptions{
provider: "openai",
target: "support.voice",
transport: "webrtc",
sdpSource: "v=0\r\noffer\r\n",
outputFormat: "raw",
Expand Down
86 changes: 68 additions & 18 deletions control-plane/internal/handlers/sessions.go
Original file line number Diff line number Diff line change
Expand Up @@ -77,29 +77,48 @@ func StartSessionHandler(store storage.StorageProvider) gin.HandlerFunc {
return
}

turnDetection, err := types.ParseSessionTurnDetection(capability.Provider, capability.Transport, definition.TurnDetection)
if err != nil {
c.JSON(http.StatusBadRequest, gin.H{"error": err.Error()})
return
}

sessionID := "sess_" + time.Now().UTC().Format("20060102_150405") + "_" + shortRandom()
model := firstNonEmptySession(req.Model, definition.Model)
voice := firstNonEmptySession(req.Voice, definition.Voice)
// The offer endpoint is stateless. Carry the registered target in its
// returned URL so the next request can resolve and validate its config.
offerQuery := url.Values{
"target": {nodeID + "." + sessionName},
"provider": {capability.Provider},
"transport": {capability.Transport},
}
if model != "" {
offerQuery.Set("model", model)
}
if voice != "" {
offerQuery.Set("voice", voice)
}
c.JSON(http.StatusCreated, gin.H{
"session_id": sessionID,
"target": nodeID + "." + sessionName,
"provider": capability.Provider,
"transport": capability.Transport,
"model": model,
"voice": voice,
"modalities": definition.Modalities,
"tags": definition.ApprovedTags,
"tool_targets": sessionToolTargets(nodeID, definition.Tools),
"offer_url": fmt.Sprintf("/api/v1/session-instances/%s/realtime-offer", url.PathEscape(sessionID)),
"tool_url": fmt.Sprintf("/api/v1/session-instances/%s/tools/{tool}", url.PathEscape(sessionID)),
"created_at": time.Now().UTC().Format(time.RFC3339Nano),
"turn_detection": turnDetection,
"session_id": sessionID,
"target": nodeID + "." + sessionName,
"provider": capability.Provider,
"transport": capability.Transport,
"model": model,
"voice": voice,
"modalities": definition.Modalities,
"tags": definition.ApprovedTags,
"tool_targets": sessionToolTargets(nodeID, definition.Tools),
"offer_url": fmt.Sprintf("/api/v1/session-instances/%s/realtime-offer", url.PathEscape(sessionID)) + "?" + offerQuery.Encode(),
"tool_url": fmt.Sprintf("/api/v1/session-instances/%s/tools/{tool}", url.PathEscape(sessionID)),
"created_at": time.Now().UTC().Format(time.RFC3339Nano),
})
}
}

func SessionRealtimeOfferHandler(store storage.StorageProvider) gin.HandlerFunc {
return func(c *gin.Context) {
_ = store
provider := strings.TrimSpace(c.Query("provider"))
transport := strings.TrimSpace(c.Query("transport"))
if provider == "" || transport == "" {
Expand All @@ -120,6 +139,33 @@ func SessionRealtimeOfferHandler(store storage.StorageProvider) gin.HandlerFunc
c.JSON(http.StatusBadRequest, gin.H{"error": "webrtc realtime offers currently require provider=openai"})
return
}
var rawTurnDetection json.RawMessage
model, voice := c.Query("model"), c.Query("voice")
if _, supplied := c.Request.URL.Query()["target"]; supplied {
target := c.Query("target")
nodeID, sessionName, ok := splitSessionTarget(target)
if !ok {
c.JSON(http.StatusBadRequest, gin.H{"error": "session target must be <node>.<session>"})
return
}
definition, found := lookupSessionDefinition(c, store, nodeID, sessionName)
if !found {
return
}
if types.NormalizeSessionTransportValue(definition.Provider) != "openai" ||
types.NormalizeSessionTransportValue(definition.Transport) != "webrtc" {
c.JSON(http.StatusBadRequest, gin.H{"error": "registered session must use provider=openai transport=webrtc for realtime offers"})
return
}
rawTurnDetection = definition.TurnDetection
model = firstNonEmptySession(model, definition.Model)
voice = firstNonEmptySession(voice, definition.Voice)
}
turnDetection, err := types.ParseSessionTurnDetection("openai", "webrtc", rawTurnDetection)
if err != nil {
c.JSON(http.StatusBadRequest, gin.H{"error": err.Error()})
return
}
if strings.TrimSpace(os.Getenv("OPENAI_API_KEY")) == "" {
c.JSON(http.StatusBadGateway, gin.H{"error": "OPENAI_API_KEY is required for provider=openai transport=webrtc"})
return
Expand All @@ -133,8 +179,9 @@ func SessionRealtimeOfferHandler(store storage.StorageProvider) gin.HandlerFunc
c.Request.Context(),
sessionPathID(c),
string(sdp),
firstNonEmptySession(c.Query("model"), "gpt-realtime-2"),
firstNonEmptySession(c.Query("voice"), "marin"),
firstNonEmptySession(model, "gpt-realtime-2"),
firstNonEmptySession(voice, "marin"),
turnDetection,
)
if err != nil {
c.JSON(http.StatusBadGateway, gin.H{
Expand Down Expand Up @@ -269,7 +316,7 @@ func sessionToolTargets(nodeID string, tools []string) map[string]string {
return targets
}

func createOpenAIRealtimeCall(ctx context.Context, sessionID string, sdp string, model string, voice string) (string, error) {
func createOpenAIRealtimeCall(ctx context.Context, sessionID string, sdp string, model string, voice string, turnDetection *types.TurnDetection) (string, error) {
var body bytes.Buffer
writer := multipart.NewWriter(&body)
if err := writer.WriteField("sdp", sdp); err != nil {
Expand All @@ -279,8 +326,11 @@ func createOpenAIRealtimeCall(ctx context.Context, sessionID string, sdp string,
"type": "realtime",
"model": model,
"instructions": "You are a realtime voice front end for an AgentField session. Use registered tools to route agent work through the AgentField control plane.",
"audio": map[string]interface{}{"output": map[string]interface{}{"voice": voice}},
"tool_choice": "auto",
"audio": map[string]interface{}{
"input": map[string]interface{}{"turn_detection": turnDetection},
"output": map[string]interface{}{"voice": voice},
},
"tool_choice": "auto",
}
sessionBytes, _ := json.Marshal(sessionConfig)
if err := writer.WriteField("session", string(sessionBytes)); err != nil {
Expand Down
161 changes: 161 additions & 0 deletions control-plane/internal/handlers/sessions_test.go
Original file line number Diff line number Diff line change
Expand Up @@ -6,6 +6,7 @@ import (
"io"
"net/http"
"net/http/httptest"
"net/url"
"strings"
"testing"
"time"
Expand Down Expand Up @@ -277,6 +278,10 @@ func TestSessionRealtimeOfferHandlerCallsRealtimeProvider(t *testing.T) {
require.Len(t, gotSafetyID, 32)
require.Contains(t, gotSession, `"model":"gpt-test"`)
require.Contains(t, gotSession, `"voice":"cedar"`)
var config map[string]interface{}
require.NoError(t, json.Unmarshal([]byte(gotSession), &config))
input := config["audio"].(map[string]interface{})["input"].(map[string]interface{})
require.Equal(t, true, input["turn_detection"].(map[string]interface{})["interrupt_response"])
}

func TestSessionRealtimeOfferHandlerSurfacesProviderErrors(t *testing.T) {
Expand Down Expand Up @@ -403,3 +408,159 @@ func sessionTestAgent() *types.AgentNode {
}},
}
}

// Exercise the complete metadata -> start -> offer -> provider boundary, rather
// than just checking that a new field exists in the registration response.
func TestSessionTurnDetectionReachesProvider(t *testing.T) {
gin.SetMode(gin.TestMode)
t.Setenv("OPENAI_API_KEY", "test-key")
original := http.DefaultClient.Transport
t.Cleanup(func() { http.DefaultClient.Transport = original })
for _, tc := range []struct{ name, config, expected string }{
{"legacy defaults", "", `{"type":"server_vad","threshold":0.5,"prefix_padding_ms":300,"silence_duration_ms":500,"create_response":true,"interrupt_response":true}`},
{"custom server", `{"type":"server_vad","threshold":0,"prefix_padding_ms":0,"silence_duration_ms":750,"create_response":false,"interrupt_response":false}`, `{"type":"server_vad","threshold":0,"prefix_padding_ms":0,"silence_duration_ms":750,"create_response":false,"interrupt_response":false}`},
{"semantic", `{"type":"semantic_vad","eagerness":"low"}`, `{"type":"semantic_vad","eagerness":"low","create_response":true,"interrupt_response":true}`},
} {
t.Run(tc.name, func(t *testing.T) {
agent := sessionTestAgent()
raw := agent.Metadata.Custom["sessions"].([]interface{})[0].(map[string]interface{})
raw["model"], raw["voice"] = "gpt-vad-test", "cedar"
if tc.config != "" {
raw["turn_detection"] = json.RawMessage(tc.config)
}
store := &nodeRESTStorageStub{agent: agent}
router := gin.New()
router.POST("/api/v1/session-targets/:target/start", StartSessionHandler(store))
router.POST("/api/v1/session-instances/:session_id/realtime-offer", SessionRealtimeOfferHandler(store))
start := httptest.NewRecorder()
router.ServeHTTP(start, httptest.NewRequest(http.MethodPost, "/api/v1/session-targets/support.voice/start", strings.NewReader(`{}`)))
require.Equal(t, http.StatusCreated, start.Code, start.Body.String())
var result struct {
OfferURL string `json:"offer_url"`
TurnDetection json.RawMessage `json:"turn_detection"`
}
require.NoError(t, json.Unmarshal(start.Body.Bytes(), &result))
require.JSONEq(t, tc.expected, string(result.TurnDetection))
calls := 0
http.DefaultClient.Transport = roundTripFunc(func(req *http.Request) (*http.Response, error) {
calls++
require.NoError(t, req.ParseMultipartForm(1<<20))
require.Equal(t, "v=0\r\noffer\r\n", req.FormValue("sdp"))
var config struct {
Model string `json:"model"`
Audio struct {
Output struct {
Voice string `json:"voice"`
} `json:"output"`
Input struct {
TurnDetection json.RawMessage `json:"turn_detection"`
} `json:"input"`
} `json:"audio"`
}
require.NoError(t, json.Unmarshal([]byte(req.FormValue("session")), &config))
require.JSONEq(t, tc.expected, string(config.Audio.Input.TurnDetection))
require.Equal(t, "gpt-vad-test", config.Model)
require.Equal(t, "cedar", config.Audio.Output.Voice)
return &http.Response{StatusCode: http.StatusOK, Header: make(http.Header), Body: io.NopCloser(strings.NewReader("answer"))}, nil
})
if tc.name == "semantic" {
// CLI offers identify the target without repeating model/voice.
offerURL, err := url.Parse(result.OfferURL)
require.NoError(t, err)
query := offerURL.Query()
query.Del("model")
query.Del("voice")
offerURL.RawQuery = query.Encode()
result.OfferURL = offerURL.String()
}
offer := httptest.NewRecorder()
router.ServeHTTP(offer, httptest.NewRequest(http.MethodPost, result.OfferURL, strings.NewReader("v=0\r\noffer\r\n")))
require.Equal(t, http.StatusOK, offer.Code, offer.Body.String())
require.Equal(t, 1, calls)
})
}
}

func TestStartSessionRejectsInvalidTurnDetection(t *testing.T) {
for _, config := range []string{
`{}`, `{"Type":"server_vad"}`, `{"type":"semantic_vad","eagerness":""}`, `{"type":"client_vad"}`, `{"type":"server_vad","threshold":2}`,
`{"type":"server_vad","silence_duration_ms":-1}`, `{"type":"server_vad","prefix_padding_ms":1.5}`,
`{"type":"server_vad","create_response":"false"}`, `{"type":"server_vad","interrupt_response":null}`,
`{"type":"semantic_vad","threshold":0.5}`, `{"type":"server_vad","eagerness":"low"}`,
`{"type":"semantic_vad","eagerness":"urgent"}`, `{"type":"server_vad","unknown":true}`, `[]`,
} {
t.Run(config, func(t *testing.T) {
agent := sessionTestAgent()
agent.Metadata.Custom["sessions"].([]interface{})[0].(map[string]interface{})["turn_detection"] = json.RawMessage(config)
router := gin.New()
router.POST("/:target/start", StartSessionHandler(&nodeRESTStorageStub{agent: agent}))
rec := httptest.NewRecorder()
router.ServeHTTP(rec, httptest.NewRequest(http.MethodPost, "/support.voice/start", strings.NewReader(`{}`)))
require.Equal(t, http.StatusBadRequest, rec.Code, rec.Body.String())
require.Contains(t, rec.Body.String(), "turn_detection")
})
}
}

func TestSessionOfferRevalidatesRegisteredTurnDetection(t *testing.T) {
t.Setenv("OPENAI_API_KEY", "") // Invalid config must fail before credentials/upstream.
agent := sessionTestAgent()
agent.Metadata.Custom["sessions"].([]interface{})[0].(map[string]interface{})["turn_detection"] = map[string]interface{}{
"type": "semantic_vad", "silence_duration_ms": 500,
}
router := gin.New()
router.POST("/:target/realtime-offer", SessionRealtimeOfferHandler(&nodeRESTStorageStub{agent: agent}))
rec := httptest.NewRecorder()
router.ServeHTTP(rec, httptest.NewRequest(http.MethodPost,
"/sess-1/realtime-offer?provider=openai&transport=webrtc&target=support.voice", strings.NewReader("v=0")))
require.Equal(t, http.StatusBadRequest, rec.Code)
require.Contains(t, rec.Body.String(), "turn_detection")
}

func TestStartSessionRejectsTurnDetectionForOpenRouter(t *testing.T) {
agent := sessionTestAgent()
raw := agent.Metadata.Custom["sessions"].([]interface{})[0].(map[string]interface{})
raw["provider"], raw["transport"] = "openrouter", "audio_turns"
raw["turn_detection"] = map[string]interface{}{"type": "server_vad"}
router := gin.New()
router.POST("/:target/start", StartSessionHandler(&nodeRESTStorageStub{agent: agent}))
rec := httptest.NewRecorder()
router.ServeHTTP(rec, httptest.NewRequest(http.MethodPost, "/support.voice/start", strings.NewReader(`{}`)))
require.Equal(t, http.StatusBadRequest, rec.Code)
require.Contains(t, rec.Body.String(), "turn_detection requires")
}

func TestSessionOfferRejectsInvalidRegisteredTargets(t *testing.T) {
gin.SetMode(gin.TestMode)
t.Setenv("OPENAI_API_KEY", "test-key")
original := http.DefaultClient.Transport
t.Cleanup(func() { http.DefaultClient.Transport = original })
http.DefaultClient.Transport = roundTripFunc(func(req *http.Request) (*http.Response, error) {
t.Fatal("invalid target must not reach the provider")
return nil, nil
})
for _, tc := range []struct {
name, target, provider, transport string
status int
message string
}{
{"empty target", "", "openai", "webrtc", http.StatusBadRequest, "session target must be"},
{"malformed target", "support", "openai", "webrtc", http.StatusBadRequest, "session target must be"},
{"missing session", "support.missing", "openai", "webrtc", http.StatusNotFound, "session not registered"},
{"wrong provider", "support.voice", "openrouter", "audio_turns", http.StatusBadRequest, "registered session must use"},
{"wrong transport", "support.voice", "openai", "websocket", http.StatusBadRequest, "registered session must use"},
} {
t.Run(tc.name, func(t *testing.T) {
agent := sessionTestAgent()
raw := agent.Metadata.Custom["sessions"].([]interface{})[0].(map[string]interface{})
raw["provider"], raw["transport"] = tc.provider, tc.transport
router := gin.New()
router.POST("/:session_id/realtime-offer", SessionRealtimeOfferHandler(&nodeRESTStorageStub{agent: agent}))
rec := httptest.NewRecorder()
router.ServeHTTP(rec, httptest.NewRequest(http.MethodPost,
"/sess-1/realtime-offer?provider=openai&transport=webrtc&target="+url.QueryEscape(tc.target), strings.NewReader("v=0")))
require.Equal(t, tc.status, rec.Code, rec.Body.String())
require.Contains(t, rec.Body.String(), tc.message)
})
}
}
Loading
Loading