From 81cf27874983ff5bea419361a908c2bcde65befa Mon Sep 17 00:00:00 2001 From: Vincent Koc Date: Sat, 22 Aug 2026 18:20:11 +0900 Subject: [PATCH] fix(talk): bind legacy iOS output cancellation Co-authored-by: Zhilong Zheng --- .../Voice/RealtimeTalkRelaySession.swift | 16 ++- .../Tests/RealtimeTalkRelaySessionTests.swift | 102 ++++++++++++++++++ docs/plugins/sdk-migration.md | 2 + 3 files changed, 118 insertions(+), 2 deletions(-) diff --git a/apps/ios/Sources/Voice/RealtimeTalkRelaySession.swift b/apps/ios/Sources/Voice/RealtimeTalkRelaySession.swift index fe6623a71a5e..f1420299b57f 100644 --- a/apps/ios/Sources/Voice/RealtimeTalkRelaySession.swift +++ b/apps/ios/Sources/Voice/RealtimeTalkRelaySession.swift @@ -166,6 +166,7 @@ final class RealtimeTalkRelaySession { private var outputContinuation: AsyncThrowingStream.Continuation? private var outputIdleTask: Task? private var outputSessionId = 0 + private var activeOutputTurnId: String? private var pendingOutputChunks: [Data] = [] private var pendingOutputDone = false private var pendingPlaybackMarks: [String] = [] @@ -334,18 +335,25 @@ final class RealtimeTalkRelaySession { _ = try? await transport.request("talk.session.close", json, 8) } - func cancelOutput(reason: String = "user") { + @discardableResult + func cancelOutput(reason: String = "user") -> Bool { + let turnId = self.activeOutputTurnId self.stopOutputPlayback() - guard let relaySessionId, let startupTransport else { return } + guard let relaySessionId, + let startupTransport, + let turnId + else { return false } Task { [startupTransport] in let payload: [String: Any] = [ "sessionId": relaySessionId, "reason": reason, + "turnId": turnId, ] let data = try? JSONSerialization.data(withJSONObject: payload) let json = data.flatMap { String(data: $0, encoding: .utf8) } _ = try? await startupTransport.request("talk.session.cancelOutput", json, 8) } + return true } private func createRelaySession() async throws -> TalkSessionCreateResult { @@ -424,6 +432,8 @@ final class RealtimeTalkRelaySession { guard let base64 = payload["audioBase64"]?.stringValue, let data = Data(base64Encoded: base64) else { return } + self.activeOutputTurnId = + self.nonEmpty(payload["talkEvent"]?.dictionaryValue?["turnId"]?.stringValue) self.recordOutputAudioChunk(byteCount: data.count) self.markOutputAudioStarted(byteCount: data.count, nowMs: ProcessInfo.processInfo.systemUptime * 1000) self.onSpeakingChanged(true) @@ -895,6 +905,7 @@ final class RealtimeTalkRelaySession { self.outputIdleTask = nil } self.isOutputPlaying = false + self.activeOutputTurnId = nil self.outputStartedAtMs = nil self.outputPlaybackExpectedEndMs = 0 self.outputEnvelope?.cancel() @@ -952,6 +963,7 @@ final class RealtimeTalkRelaySession { self.pendingOutputDone = false _ = self.pcmPlayer.stop() self.isOutputPlaying = false + self.activeOutputTurnId = nil self.outputStartedAtMs = nil self.outputPlaybackExpectedEndMs = 0 self.outputEnvelope?.cancel() diff --git a/apps/ios/Tests/RealtimeTalkRelaySessionTests.swift b/apps/ios/Tests/RealtimeTalkRelaySessionTests.swift index 397eeeee7ce0..daec6569973b 100644 --- a/apps/ios/Tests/RealtimeTalkRelaySessionTests.swift +++ b/apps/ios/Tests/RealtimeTalkRelaySessionTests.swift @@ -15,6 +15,22 @@ private final class UnusedPCMStreamingAudioPlayer: PCMStreamingAudioPlaying { } } +@MainActor +private final class DrainingPCMStreamingAudioPlayer: PCMStreamingAudioPlaying { + func play(stream: AsyncThrowingStream, sampleRate: Double) async -> StreamingPlaybackResult { + do { + for try await _ in stream {} + return StreamingPlaybackResult(finished: true, interruptedAt: nil) + } catch { + return StreamingPlaybackResult(finished: false, interruptedAt: nil) + } + } + + func stop() -> Double? { + nil + } +} + private actor RealtimeRelayStartupBarrier { private var entered = false private var enteredWaiter: CheckedContinuation? @@ -59,6 +75,75 @@ private actor RealtimeRelayStartupRequestLog { @MainActor struct RealtimeTalkRelaySessionTests { + @Test func `delayed cancellation remains bound to the output turn it stopped`() async throws { + let requestStarted = RealtimeRelayStartupBarrier() + let requestRecorded = RealtimeRelayStartupBarrier() + let requests = RealtimeRelayStartupRequestLog() + let transport = RealtimeTalkRelaySession.StartupTransport( + subscribeServerEvents: { _ in AsyncStream { $0.finish() } }, + request: { method, paramsJSON, _ in + if method == "talk.session.cancelOutput" { + await requestStarted.suspend() + } + await requests.record(method: method, paramsJSON: paramsJSON) + if method == "talk.session.cancelOutput" { + await requestRecorded.suspend() + } + return Data("{\"ok\":true}".utf8) + }) + let session = RealtimeTalkRelaySession( + gateway: GatewayNodeSession(), + options: .init(sessionKey: "main", provider: "xai", model: nil, voice: nil), + pcmPlayer: DrainingPCMStreamingAudioPlayer(), + onStatus: { _ in }, + onSpeakingChanged: { _ in }, + startupTransport: transport) + defer { session.stop() } + session._test_setRelaySessionId("relay-1") + + await session._test_handleGatewayEvent(Self.audioEvent(turnId: "turn-a")) + session.cancelOutput(reason: "barge-in") + await requestStarted.waitUntilEntered() + + await session._test_handleGatewayEvent(Self.audioEvent(turnId: "turn-b")) + await requestStarted.release() + await requestRecorded.waitUntilEntered() + await requestRecorded.release() + + let request = try #require(await requests.snapshot().first) + let paramsData = try #require(request.paramsJSON?.data(using: .utf8)) + let params = try #require(JSONSerialization.jsonObject(with: paramsData) as? [String: String]) + #expect(request.method == "talk.session.cancelOutput") + #expect(params["turnId"] == "turn-a") + } + + @Test func `id less replacement cannot reuse the prior output turn`() async { + let requests = RealtimeRelayStartupRequestLog() + let transport = RealtimeTalkRelaySession.StartupTransport( + subscribeServerEvents: { _ in AsyncStream { $0.finish() } }, + request: { method, paramsJSON, _ in + await requests.record(method: method, paramsJSON: paramsJSON) + return Data("{\"ok\":true}".utf8) + }) + let session = RealtimeTalkRelaySession( + gateway: GatewayNodeSession(), + options: .init(sessionKey: "main", provider: "xai", model: nil, voice: nil), + pcmPlayer: DrainingPCMStreamingAudioPlayer(), + onStatus: { _ in }, + onSpeakingChanged: { _ in }, + startupTransport: transport) + defer { session.stop() } + session._test_setRelaySessionId("relay-1") + + await session._test_handleGatewayEvent(Self.audioEvent(turnId: "turn-a")) + await session._test_handleGatewayEvent(Self.audioEvent(turnId: nil)) + + #expect(session._test_isOutputPlaying()) + #expect(!session.cancelOutput(reason: "barge-in")) + #expect(!session._test_isOutputPlaying()) + #expect(await requests.snapshot().isEmpty) + } + @Test func `output playback finish clears barge in start time`() { var speakingStates: [Bool] = [] let session = RealtimeTalkRelaySession( @@ -343,4 +428,21 @@ struct RealtimeTalkRelaySessionTests { #expect(await requests.snapshot().isEmpty) } + + private static func audioEvent(turnId: String?) -> EventFrame { + var payload: [String: Any] = [ + "relaySessionId": "relay-1", + "type": "audio", + "audioBase64": Data([0x01, 0x02]).base64EncodedString(), + ] + if let turnId { + payload["talkEvent"] = ["turnId": turnId] + } + return EventFrame( + type: "event", + event: "talk.event", + payload: AnyCodable(payload), + seq: nil, + stateversion: nil) + } } diff --git a/docs/plugins/sdk-migration.md b/docs/plugins/sdk-migration.md index 2eb2347b9cb0..a428329d2f91 100644 --- a/docs/plugins/sdk-migration.md +++ b/docs/plugins/sdk-migration.md @@ -981,6 +981,8 @@ await gateway.request("talk.session.create", { sessionKey: "main", }); await gateway.request("talk.session.appendAudio", { sessionId, audioBase64 }); +// Capture this before stopping playback from the active output `talk.event`. +const turnId = activeOutputTalkEvent.talkEvent.turnId; await gateway.request("talk.session.cancelOutput", { sessionId, turnId, reason: "barge-in" }); await gateway.request("talk.session.submitToolResult", { sessionId,