MockSSE.swift 7.3 KB

123456789101112131415161718192021222324252627282930313233343536373839404142434445464748495051525354555657585960616263646566676869707172737475767778798081828384858687888990919293949596979899100101102103104105106107108109110111112113114115116117118119120121122123124125126127128129130131132133134135136137138139140141142143144145146147148149150151152153154155156157158159160161162163164165166167168169170171172173174175176177178179180181182183184185186187188189190191192193194195196197198199200201202203204205206207208209210211212213214215216217218219220221222223224225226227
  1. //===----------------------------------------------------------------------===//
  2. //
  3. // This source file is part of the Foundation Models open source project.
  4. //
  5. // Copyright © 2024-2027 Apple Inc. and the Foundation Models project authors.
  6. //
  7. // Licensed under the Apache License v2.0
  8. //
  9. // See LICENSE.txt for license information
  10. //
  11. //===----------------------------------------------------------------------===//
  12. import Foundation
  13. #if canImport(FoundationNetworking)
  14. import FoundationNetworking
  15. #endif
  16. // Intercepts requests to "mock-llm.test" and returns canned SSE responses,
  17. // letting us test the full ChatCompletionsLanguageModel pipeline without a network.
  18. final class MockSSEProtocol: URLProtocol, @unchecked Sendable {
  19. nonisolated(unsafe) static var handler: ((URLRequest) -> (statusCode: Int, data: Data))?
  20. nonisolated(unsafe) static var lastRequest: URLRequest?
  21. static func reset() {
  22. handler = nil
  23. lastRequest = nil
  24. }
  25. override class func canInit(with request: URLRequest) -> Bool {
  26. request.url?.host == "mock-llm.test"
  27. }
  28. override class func canonicalRequest(for request: URLRequest) -> URLRequest {
  29. request
  30. }
  31. override func startLoading() {
  32. // URLSession converts httpBody to httpBodyStream before passing to protocol handlers;
  33. // drain it back into httpBody so handlers and lastRequest see consistent data.
  34. var normalizedRequest = request
  35. if let stream = request.httpBodyStream {
  36. stream.open()
  37. var bodyData = Data()
  38. let buffer = UnsafeMutablePointer<UInt8>.allocate(capacity: 4096)
  39. defer { buffer.deallocate() }
  40. while stream.hasBytesAvailable {
  41. let n = stream.read(buffer, maxLength: 4096)
  42. if n > 0 { bodyData.append(buffer, count: n) }
  43. }
  44. stream.close()
  45. normalizedRequest.httpBody = bodyData
  46. }
  47. Self.lastRequest = normalizedRequest
  48. let (statusCode, data) = Self.handler?(normalizedRequest) ?? (200, Data())
  49. let response = HTTPURLResponse(
  50. url: request.url!,
  51. statusCode: statusCode,
  52. httpVersion: "HTTP/1.1",
  53. headerFields: ["Content-Type": "text/event-stream"],
  54. )!
  55. client?.urlProtocol(self, didReceive: response, cacheStoragePolicy: .notAllowed)
  56. client?.urlProtocol(self, didLoad: data)
  57. client?.urlProtocolDidFinishLoading(self)
  58. }
  59. override func stopLoading() {}
  60. }
  61. enum MockSSE {
  62. static func text(_ chunks: String...) -> Data {
  63. var lines = [String]()
  64. for chunk in chunks {
  65. let escaped = jsonEscape(chunk)
  66. lines.append(
  67. #"data: {"id":"1","model":"mock","choices":[{"delta":{"content":"\#(escaped)"}}]}"#
  68. )
  69. lines.append("")
  70. }
  71. lines.append("data: [DONE]")
  72. lines.append("")
  73. return Data(lines.joined(separator: "\n").utf8)
  74. }
  75. // A single SSE chunk with optional text content, optional reasoning
  76. // content, and optional usage snapshot.
  77. struct Chunk {
  78. var text: String? = nil
  79. var reasoning: String? = nil
  80. var usage: Usage? = nil
  81. struct Usage {
  82. var promptTokens: Int
  83. var completionTokens: Int
  84. var cachedTokens: Int? = nil
  85. var reasoningTokens: Int? = nil
  86. }
  87. }
  88. // Builds an SSE response from a sequence of chunks, each carrying optional
  89. // text and/or a usage snapshot. Use this for both trailing-usage streams
  90. // (where the final chunk has only `usage`) and per-chunk-usage streams
  91. // (where every chunk has both `text` and `usage`).
  92. static func chunks(_ chunks: [Chunk]) -> Data {
  93. var lines = [String]()
  94. for chunk in chunks {
  95. lines.append("data: " + chunkJSON(chunk))
  96. lines.append("")
  97. }
  98. lines.append("data: [DONE]")
  99. lines.append("")
  100. return Data(lines.joined(separator: "\n").utf8)
  101. }
  102. private static func chunkJSON(_ chunk: Chunk) -> String {
  103. var deltaFields = [String]()
  104. if let text = chunk.text {
  105. deltaFields.append(#""content":"\#(jsonEscape(text))""#)
  106. }
  107. if let reasoning = chunk.reasoning {
  108. deltaFields.append(#""reasoning_content":"\#(jsonEscape(reasoning))""#)
  109. }
  110. let choices: String
  111. if deltaFields.isEmpty {
  112. choices = "[]"
  113. } else {
  114. let delta = "{" + deltaFields.joined(separator: ",") + "}"
  115. choices = #"[{"delta":\#(delta)}]"#
  116. }
  117. guard let usage = chunk.usage else {
  118. return #"{"id":"1","model":"mock","choices":\#(choices)}"#
  119. }
  120. var usageFields = [
  121. #""prompt_tokens":\#(usage.promptTokens)"#,
  122. #""completion_tokens":\#(usage.completionTokens)"#,
  123. #""total_tokens":\#(usage.promptTokens + usage.completionTokens)"#
  124. ]
  125. if let cachedTokens = usage.cachedTokens {
  126. usageFields.append(
  127. #""prompt_tokens_details":{"cached_tokens":\#(cachedTokens)}"#
  128. )
  129. }
  130. if let reasoningTokens = usage.reasoningTokens {
  131. usageFields.append(
  132. #""completion_tokens_details":{"reasoning_tokens":\#(reasoningTokens)}"#
  133. )
  134. }
  135. let usageJSON = "{" + usageFields.joined(separator: ",") + "}"
  136. return #"{"id":"1","model":"mock","choices":\#(choices),"usage":\#(usageJSON)}"#
  137. }
  138. static func toolCall(
  139. id: String,
  140. name: String,
  141. argumentChunks: [String]
  142. ) -> Data {
  143. var lines = [String]()
  144. lines.append(
  145. #"data: {"id":"1","model":"mock","choices":[{"delta":{"tool_calls":[{"index":0,"id":"\#(id)","type":"function","function":{"name":"\#(name)","arguments":""}}]}}]}"#
  146. )
  147. lines.append("")
  148. for chunk in argumentChunks {
  149. let escaped = jsonEscape(chunk)
  150. lines.append(
  151. #"data: {"id":"1","model":"mock","choices":[{"delta":{"tool_calls":[{"index":0,"function":{"arguments":"\#(escaped)"}}]}}]}"#
  152. )
  153. lines.append("")
  154. }
  155. lines.append("data: [DONE]")
  156. lines.append("")
  157. return Data(lines.joined(separator: "\n").utf8)
  158. }
  159. static func parallelToolCalls(
  160. _ calls: [(id: String, name: String, arguments: String)]
  161. ) -> Data {
  162. var lines = [String]()
  163. for (index, call) in calls.enumerated() {
  164. lines.append(
  165. #"data: {"id":"1","model":"mock","choices":[{"delta":{"tool_calls":[{"index":\#(index),"id":"\#(call.id)","type":"function","function":{"name":"\#(call.name)","arguments":"\#(jsonEscape(call.arguments))"}}]}}]}"#
  166. )
  167. lines.append("")
  168. }
  169. lines.append("data: [DONE]")
  170. lines.append("")
  171. return Data(lines.joined(separator: "\n").utf8)
  172. }
  173. static func apiError(message: String) -> Data {
  174. let escaped = jsonEscape(message)
  175. return Data(
  176. [
  177. #"data: {"error":{"message":"\#(escaped)","type":"server_error"}}"#,
  178. ""
  179. ].joined(separator: "\n").utf8
  180. )
  181. }
  182. static func toolCallThenText(
  183. toolCallData: Data,
  184. textResponse: String
  185. ) -> (URLRequest) -> (statusCode: Int, data: Data) {
  186. { request in
  187. let body = request.httpBody.flatMap {
  188. try? JSONSerialization.jsonObject(with: $0) as? [String: Any]
  189. }
  190. let messages = body?["messages"] as? [[String: Any]]
  191. let hasToolOutput =
  192. messages?.contains {
  193. $0["role"] as? String == "tool"
  194. } ?? false
  195. if hasToolOutput {
  196. return (200, MockSSE.text(textResponse))
  197. } else {
  198. return (200, toolCallData)
  199. }
  200. }
  201. }
  202. private static func jsonEscape(_ string: String) -> String {
  203. string
  204. .replacingOccurrences(of: "\\", with: "\\\\")
  205. .replacingOccurrences(of: "\"", with: "\\\"")
  206. .replacingOccurrences(of: "\n", with: "\\n")
  207. }
  208. }