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
52 changes: 42 additions & 10 deletions swift/Sources/Tools/llm-server/ChatHandler.swift
Original file line number Diff line number Diff line change
Expand Up @@ -7,6 +7,7 @@ import CoreAILMCommon
import CoreAILanguageModels
import CoreAIShared
import Foundation
import HTTPTypes
import Hummingbird
import NIOCore
import NIOFoundationCompat
Expand Down Expand Up @@ -42,7 +43,9 @@ func startServer(state: ServerState, port: Int) async throws {
}

router.post("/v1/chat/completions") { request, _ in
try await handleChatCompletionsRoute(request: request, state: state)
let sessionID =
HTTPField.Name("X-Session-ID").flatMap { request.headers[$0] } ?? "default"
return try await handleChatCompletionsRoute(request: request, state: state, sessionID: sessionID)
}

router.post("/v1") { request, _ in
Expand All @@ -53,6 +56,16 @@ func startServer(state: ServerState, port: Int) async throws {
try await handleCompletionsRoute(request: request, state: state)
}

router.get("/v1/stats") { _, _ in
let stats = state.statsSnapshot()
let data = try JSONEncoder().encode(stats)
return Response(
status: .ok,
headers: [.contentType: "application/json"],
body: .init(byteBuffer: ByteBuffer(data: data))
)
}

let app = Application(
router: router,
configuration: .init(
Expand All @@ -74,7 +87,9 @@ private func handleAutoRoute(request: Request, state: ServerState) async throws
}
}

private func handleChatCompletionsFromBody(body: ByteBuffer, state: ServerState) async throws -> Response {
private func handleChatCompletionsFromBody(body: ByteBuffer, state: ServerState, sessionID: String? = nil) async throws
-> Response
{
guard state.tryAcquire() else {
let err = ErrorResponse(error: .init(message: "Server is busy.", type: "server_error", code: "busy"))
let data = try JSONEncoder().encode(err)
Expand All @@ -96,9 +111,10 @@ private func handleChatCompletionsFromBody(body: ByteBuffer, state: ServerState)
do {
let shouldStream = chatRequest.stream ?? false
if shouldStream {
return try await handleStreamingRequest(chatRequest: chatRequest, state: state)
return try await handleStreamingRequest(chatRequest: chatRequest, state: state, sessionID: sessionID)
} else {
let response = try await handleNonStreamingRequest(chatRequest: chatRequest, state: state)
let response = try await handleNonStreamingRequest(
chatRequest: chatRequest, state: state, sessionID: sessionID)
state.release()
return response
}
Expand All @@ -122,14 +138,17 @@ private func handleChatCompletionsFromBody(body: ByteBuffer, state: ServerState)

// MARK: - Route Handler

private func handleChatCompletionsRoute(request: Request, state: ServerState) async throws -> Response {
private func handleChatCompletionsRoute(request: Request, state: ServerState, sessionID: String? = nil) async throws
-> Response
{
let body = try await request.body.collect(upTo: 10 * 1024 * 1024)
return try await handleChatCompletionsFromBody(body: body, state: state)
return try await handleChatCompletionsFromBody(body: body, state: state, sessionID: sessionID)
}

// MARK: - Non-Streaming

private func handleNonStreamingRequest(chatRequest: ChatCompletionRequest, state: ServerState) async throws -> Response
private func handleNonStreamingRequest(chatRequest: ChatCompletionRequest, state: ServerState, sessionID: String? = nil)
async throws -> Response
{
let requestMaxTokens = chatRequest.maxCompletionTokens ?? chatRequest.maxTokens ?? state.config.defaultMaxTokens
guard requestMaxTokens > 0 else {
Expand Down Expand Up @@ -158,7 +177,11 @@ private func handleNonStreamingRequest(chatRequest: ChatCompletionRequest, state
"[\(requestID)] messages: \(chatRequest.messages.count), tokens: \(promptTokens.count), max_tokens: \(requestMaxTokens)",
component: "Server")

try await state.engine.reset()
let promptTokensInt32 = promptTokens.map { Int32($0) }
let prefixReused = await state.prepareForRequest(sessionID: sessionID, promptTokens: promptTokensInt32)
if prefixReused > 0 {
CLILogger.log("[\(requestID)] prefix reuse: \(prefixReused) tokens cached", component: "Server")
}
let t0 = SuspendingClock().now

let strategy: any DecodingStrategy
Expand Down Expand Up @@ -198,6 +221,7 @@ private func handleNonStreamingRequest(chatRequest: ChatCompletionRequest, state
state.stats.record(
promptTokens: promptTokens.count, genTokens: genTokenCount, promptSeconds: 0, genSeconds: seconds,
totalSeconds: seconds)
state.recordPromptTokens(promptTokensInt32)

let response = ChatCompletionResponse(
id: requestID,
Expand Down Expand Up @@ -228,7 +252,9 @@ private func handleNonStreamingRequest(chatRequest: ChatCompletionRequest, state

// MARK: - Streaming (SSE)

private func handleStreamingRequest(chatRequest: ChatCompletionRequest, state: ServerState) async throws -> Response {
private func handleStreamingRequest(chatRequest: ChatCompletionRequest, state: ServerState, sessionID: String? = nil)
async throws -> Response
{
let requestMaxTokens = chatRequest.maxCompletionTokens ?? chatRequest.maxTokens ?? state.config.defaultMaxTokens
guard requestMaxTokens > 0 else {
throw ServerError.badRequest("max_tokens must be positive")
Expand Down Expand Up @@ -256,10 +282,15 @@ private func handleStreamingRequest(chatRequest: ChatCompletionRequest, state: S
"[\(requestID)] stream, messages: \(chatRequest.messages.count), tokens: \(promptTokens.count), max_tokens: \(requestMaxTokens)",
component: "Server")

let promptTokensInt32 = promptTokens.map { Int32($0) }
let prefixReused = await state.prepareForRequest(sessionID: sessionID, promptTokens: promptTokensInt32)
if prefixReused > 0 {
CLILogger.log("[\(requestID)] prefix reuse: \(prefixReused) tokens cached", component: "Server")
}

let responseBody = ResponseBody { writer in
defer { state.release() }
do {
try await state.engine.reset()
let encoder = JSONEncoder()
let genStart = SuspendingClock().now

Expand Down Expand Up @@ -341,6 +372,7 @@ private func handleStreamingRequest(chatRequest: ChatCompletionRequest, state: S
state.stats.record(
promptTokens: promptTokens.count, genTokens: tokenCount, promptSeconds: 0, genSeconds: seconds,
totalSeconds: seconds)
state.recordPromptTokens(promptTokensInt32)

try await writer.finish(nil)
} catch {
Expand Down
110 changes: 105 additions & 5 deletions swift/Sources/Tools/llm-server/ServerState.swift
Original file line number Diff line number Diff line change
Expand Up @@ -82,6 +82,20 @@ final class ServerStats: @unchecked Sendable {
==================================================
""")
}

func buildResponse(prefixHitRate: Double, prefixHits: Int, prefixMisses: Int) -> StatsResponse {
let s = lock.withLock { $0 }
return StatsResponse(
totalRequests: s.totalRequests,
totalPromptTokens: s.totalPromptTokens,
totalGenTokens: s.totalGenTokens,
avgPrefillTokPerSec: s.totalPromptSeconds > 0 ? Double(s.totalPromptTokens) / s.totalPromptSeconds : 0,
avgDecodeTokPerSec: s.totalGenSeconds > 0 ? Double(s.totalGenTokens) / s.totalGenSeconds : 0,
prefixHitRate: prefixHitRate,
prefixHits: prefixHits,
prefixMisses: prefixMisses
)
}
}

// MARK: - Server State
Expand All @@ -91,7 +105,15 @@ final class ServerState: @unchecked Sendable {
let tokenizer: any Tokenizer
let config: ServerConfig
let stats = ServerStats()
private let _generating = Mutex<Bool>(false)
private let _state = Mutex<InternalState>(InternalState())

private struct InternalState {
var generating: Bool = false
var lastSessionID: String? = nil
var lastPromptTokens: [Int32] = []
var prefixHits: Int = 0
var prefixMisses: Int = 0
}

init(engine: any InferenceEngine, tokenizer: any Tokenizer, config: ServerConfig) {
self.engine = engine
Expand All @@ -100,15 +122,69 @@ final class ServerState: @unchecked Sendable {
}

func tryAcquire() -> Bool {
_generating.withLock { busy in
guard !busy else { return false }
busy = true
_state.withLock { s in
guard !s.generating else { return false }
s.generating = true
return true
}
}

func release() {
_generating.withLock { $0 = false }
_state.withLock { $0.generating = false }
}

/// Prepare engine for a new request. Returns the number of prefix tokens reused.
func prepareForRequest(sessionID: String?, promptTokens: [Int32]) async -> Int {
let action = _state.withLock { s -> PrepareAction in
guard let sid = sessionID, sid == s.lastSessionID else {
s.lastSessionID = sessionID
s.lastPromptTokens = []
s.prefixMisses += 1
return .reset
}

let cached = s.lastPromptTokens
var match = 0
let limit = min(cached.count, promptTokens.count)
while match < limit && cached[match] == promptTokens[match] {
match += 1
}

if match == 0 {
s.prefixMisses += 1
return .reset
}

s.prefixHits += 1
return .reuse(prefixLength: match)
}

switch action {
case .reset:
return 0
case .reuse(let prefixLength):
return prefixLength
}
}

/// Record the tokens that were processed (call after generate completes).
func recordPromptTokens(_ tokens: [Int32]) {
_state.withLock { $0.lastPromptTokens = tokens }
}

/// Stats snapshot for /v1/stats endpoint.
func statsSnapshot() -> StatsResponse {
let s = _state.withLock { s in
(prefixHits: s.prefixHits, prefixMisses: s.prefixMisses)
}
let hitTotal = s.prefixHits + s.prefixMisses
let hitRate = hitTotal > 0 ? Double(s.prefixHits) / Double(hitTotal) : 0
return stats.buildResponse(prefixHitRate: hitRate, prefixHits: s.prefixHits, prefixMisses: s.prefixMisses)
}

private enum PrepareAction {
case reset
case reuse(prefixLength: Int)
}

func makeSamplingConfig(
Expand Down Expand Up @@ -160,3 +236,27 @@ enum ServerError: Error, LocalizedError {
}
}
}

// MARK: - Stats Response

struct StatsResponse: Codable, Sendable {
let totalRequests: Int
let totalPromptTokens: Int
let totalGenTokens: Int
let avgPrefillTokPerSec: Double
let avgDecodeTokPerSec: Double
let prefixHitRate: Double
let prefixHits: Int
let prefixMisses: Int

enum CodingKeys: String, CodingKey {
case totalRequests = "total_requests"
case totalPromptTokens = "total_prompt_tokens"
case totalGenTokens = "total_gen_tokens"
case avgPrefillTokPerSec = "avg_prefill_tok_per_sec"
case avgDecodeTokPerSec = "avg_decode_tok_per_sec"
case prefixHitRate = "prefix_hit_rate"
case prefixHits = "prefix_hits"
case prefixMisses = "prefix_misses"
}
}