Enforce deterministic gateway task policies

Make the server authoritative for thinking, model, search, tools, retry, and output budgets while preserving legacy client behavior.
This commit is contained in:
Rocky
2026-08-19 20:55:05 +08:00
parent 3edc86a9a0
commit 11ec34dacb
17 changed files with 667 additions and 44 deletions
@@ -2,6 +2,8 @@ package com.osglab.account.features.gateway.models
import io.kotest.assertions.throwables.shouldThrow
import io.kotest.core.spec.style.StringSpec
import io.kotest.matchers.shouldBe
import kotlinx.serialization.json.Json
class TextRequestPolicyTest : StringSpec({
"rejects blank and oversized input" {
@@ -54,4 +56,20 @@ class TextRequestPolicyTest : StringSpec({
TextRequestPolicy.validate(request, GatewayCapability.POLISH)
}
}
"keeps taskKind wire values stable" {
val values = mapOf(
"dictation_polish" to GatewayTaskKind.DICTATION_POLISH,
"translation" to GatewayTaskKind.TRANSLATION,
"edit_last_input" to GatewayTaskKind.EDIT_LAST_INPUT,
"ai_question" to GatewayTaskKind.AI_QUESTION,
"clipboard_transform" to GatewayTaskKind.CLIPBOARD_TRANSFORM,
"custom_skill" to GatewayTaskKind.CUSTOM_SKILL,
"agent_planning" to GatewayTaskKind.AGENT_PLANNING,
)
values.forEach { (wireValue, taskKind) ->
Json.decodeFromString<GatewayTaskKind>("\"$wireValue\"") shouldBe taskKind
}
}
})
@@ -1,14 +1,17 @@
package com.osglab.account.features.gateway.providers.deepseek
import com.osglab.account.features.gateway.models.GatewayCapability
import com.osglab.account.features.gateway.models.GatewayTaskKind
import com.osglab.account.features.gateway.models.ProviderOutput
import com.osglab.account.features.gateway.models.TextProviderRequest
import com.osglab.account.features.gateway.services.GatewayTaskPolicyResolver
import io.kotest.assertions.throwables.shouldThrow
import io.kotest.core.spec.style.StringSpec
import io.kotest.matchers.shouldBe
import io.ktor.client.HttpClient
import io.ktor.client.engine.mock.MockEngine
import io.ktor.client.engine.mock.respond
import io.ktor.client.engine.mock.toByteArray
import io.ktor.client.plugins.contentnegotiation.ContentNegotiation
import io.ktor.http.ContentType
import io.ktor.http.HttpHeaders
@@ -16,6 +19,8 @@ import io.ktor.http.HttpStatusCode
import io.ktor.http.headersOf
import io.ktor.serialization.kotlinx.json.json
import kotlinx.serialization.json.Json
import kotlinx.serialization.json.jsonObject
import kotlinx.serialization.json.jsonPrimitive
class DeepSeekClientTest : StringSpec({
"prefers provider token usage" {
@@ -79,6 +84,55 @@ class DeepSeekClientTest : StringSpec({
}
}
"serializes disabled thinking for polish without search or tools" {
var requestBody = ""
val client = client(
"""{"choices":[{"message":{"content":"ok"}}],"usage":{"prompt_tokens":10,"completion_tokens":3,"total_tokens":13}}""",
onRequest = { requestBody = it },
)
try {
KtorDeepSeekClient(client, CONFIG).complete(
request(capability = GatewayCapability.POLISH),
DISCARD_OUTPUT,
)
val payload = Json.parseToJsonElement(requestBody).jsonObject
payload.getValue("model").jsonPrimitive.content shouldBe "configured-fast-model"
payload.getValue("thinking").jsonObject.getValue("type").jsonPrimitive.content shouldBe
"disabled"
payload["reasoning_effort"] shouldBe null
payload["web_search"] shouldBe null
payload["enable_search"] shouldBe null
payload["tools"] shouldBe null
payload["tool_choice"] shouldBe null
} finally {
client.close()
}
}
"serializes explicit high-effort thinking and the reasoning model for AI" {
var requestBody = ""
val client = client(
"""{"choices":[{"message":{"content":"ok"}}],"usage":{"prompt_tokens":10,"completion_tokens":3,"total_tokens":13}}""",
onRequest = { requestBody = it },
)
try {
KtorDeepSeekClient(client, CONFIG).complete(request(), DISCARD_OUTPUT)
val payload = Json.parseToJsonElement(requestBody).jsonObject
payload.getValue("model").jsonPrimitive.content shouldBe "configured-reasoning-model"
payload.getValue("thinking").jsonObject.getValue("type").jsonPrimitive.content shouldBe
"enabled"
payload.getValue("reasoning_effort").jsonPrimitive.content shouldBe "high"
payload["web_search"] shouldBe null
payload["enable_search"] shouldBe null
payload["tools"] shouldBe null
payload["tool_choice"] shouldBe null
} finally {
client.close()
}
}
"retries one buffered empty result and returns the successful retry" {
var attempts = 0
val provider = DeepSeekProvider(
@@ -117,6 +171,22 @@ class DeepSeekClientTest : StringSpec({
attempts shouldBe 2
}
"does not retry an empty buffered polish result" {
var attempts = 0
val provider = DeepSeekProvider(
DeepSeekClient { _, _ ->
attempts += 1
throw DeepSeekEmptyResultException()
},
)
shouldThrow<DeepSeekEmptyResultException> {
provider.execute(request(capability = GatewayCapability.POLISH), DISCARD_OUTPUT)
}
attempts shouldBe 1
}
"does not retry an empty streaming result after output may have started" {
var attempts = 0
val provider = DeepSeekProvider(
@@ -235,8 +305,10 @@ class DeepSeekClientTest : StringSpec({
private fun client(
responseBody: String,
contentType: ContentType = ContentType.Application.Json,
onRequest: suspend (String) -> Unit = {},
) = HttpClient(
MockEngine {
MockEngine { request ->
onRequest(request.body.toByteArray().decodeToString())
respond(
content = responseBody,
status = HttpStatusCode.OK,
@@ -249,21 +321,30 @@ private fun client(
}
}
private fun request(capability: GatewayCapability = GatewayCapability.AI) = TextProviderRequest(
requestId = "deepseek-request",
capability = capability,
input = "hello",
context = null,
maxOutputTokens = 32,
temperature = 0.2,
stream = false,
)
private fun request(
capability: GatewayCapability = GatewayCapability.AI,
taskKind: GatewayTaskKind? = null,
): TextProviderRequest {
val executionPolicy = TASK_POLICY.resolve(capability, taskKind, 32)
return TextProviderRequest(
requestId = "deepseek-request",
capability = capability,
executionPolicy = executionPolicy,
input = "hello",
context = null,
maxOutputTokens = executionPolicy.maxOutputTokens,
temperature = 0.2,
stream = false,
)
}
private val CONFIG = DeepSeekConfig(
endpoint = "https://api.deepseek.com/v1",
apiKey = "test-key",
model = "configured-model",
model = "configured-fast-model",
reasoningModel = "configured-reasoning-model",
)
private val TASK_POLICY = GatewayTaskPolicyResolver()
private val SUCCESS_USAGE = com.osglab.account.features.gateway.models.ProviderUsage(
meter = com.osglab.account.features.gateway.models.UsageMeter.LLM_TOKEN,
units = 13,
@@ -3,10 +3,13 @@ package com.osglab.account.features.gateway.routes
import com.osglab.account.features.gateway.models.GatewayCapability
import com.osglab.account.features.gateway.models.GatewayPrincipal
import com.osglab.account.features.gateway.models.GatewayRequestSource
import com.osglab.account.features.gateway.models.GatewayTaskKind
import com.osglab.account.features.gateway.models.GatewayThinkingMode
import com.osglab.account.features.gateway.models.ProviderDescriptor
import com.osglab.account.features.gateway.models.ProviderOutput
import com.osglab.account.features.gateway.models.ProviderRequest
import com.osglab.account.features.gateway.models.ProviderUsage
import com.osglab.account.features.gateway.models.TextProviderRequest
import com.osglab.account.features.gateway.models.UsageMeter
import com.osglab.account.features.gateway.ports.CreditReservation
import com.osglab.account.features.gateway.ports.CreditReservationPort
@@ -17,7 +20,9 @@ import com.osglab.account.features.gateway.providers.GatewayProvider
import com.osglab.account.features.gateway.providers.ProviderCatalog
import com.osglab.account.features.gateway.services.GatewayService
import io.kotest.core.spec.style.StringSpec
import io.kotest.matchers.string.shouldContain
import io.kotest.matchers.shouldBe
import io.ktor.client.statement.bodyAsText
import io.ktor.client.request.header
import io.ktor.client.request.post
import io.ktor.client.request.setBody
@@ -71,6 +76,65 @@ class GatewayRequestIdTest : StringSpec({
provider.lastRequestSource shouldBe GatewayRequestSource.HOTWORD
}
}
"defaults legacy polish requests to the server-controlled low-latency policy" {
val provider = RequestIdProvider()
testApplication {
application { gatewayTestApplication(provider) }
val response = client.post("/v1/gateway/llm/polish") {
header("X-Request-ID", "legacy-polish-123")
contentType(ContentType.Application.Json)
setBody("""{"input":"hello","maxOutputTokens":4096}""")
}
response.status shouldBe HttpStatusCode.OK
provider.lastTextRequest?.executionPolicy?.taskKind shouldBe
GatewayTaskKind.DICTATION_POLISH
provider.lastTextRequest?.executionPolicy?.thinking shouldBe
GatewayThinkingMode.DISABLED
provider.lastTextRequest?.maxOutputTokens shouldBe 512
}
}
"rejects a capability and taskKind mismatch before billing or provider execution" {
val provider = RequestIdProvider()
testApplication {
application { gatewayTestApplication(provider) }
val response = client.post("/v1/gateway/llm/polish") {
header("X-Request-ID", "invalid-task-123")
contentType(ContentType.Application.Json)
setBody("""{"input":"hello","taskKind":"agent_planning"}""")
}
response.status shouldBe HttpStatusCode.BadRequest
response.bodyAsText() shouldContain """"code":"invalid_request""""
provider.calls shouldBe 0
}
}
"applies an explicit AI transform task without allowing thinking" {
val provider = RequestIdProvider()
testApplication {
application { gatewayTestApplication(provider) }
val response = client.post("/v1/gateway/llm/ai") {
header("X-Request-ID", "clipboard-task-123")
contentType(ContentType.Application.Json)
setBody("""{"input":"hello","taskKind":"clipboard_transform"}""")
}
response.status shouldBe HttpStatusCode.OK
provider.lastTextRequest?.executionPolicy?.taskKind shouldBe
GatewayTaskKind.CLIPBOARD_TRANSFORM
provider.lastTextRequest?.executionPolicy?.thinking shouldBe
GatewayThinkingMode.DISABLED
}
}
})
private fun io.ktor.server.application.Application.gatewayTestApplication(provider: RequestIdProvider) {
@@ -96,10 +160,15 @@ private class RequestIdProvider : GatewayProvider {
var calls = 0
var lastRequestId: String? = null
var lastRequestSource: GatewayRequestSource? = null
var lastTextRequest: TextProviderRequest? = null
override val descriptor = ProviderDescriptor(
id = "request-id-provider",
capabilities = setOf(GatewayCapability.AI),
capabilities = setOf(
GatewayCapability.POLISH,
GatewayCapability.AI,
GatewayCapability.AGENT,
),
streaming = true,
usageMeter = UsageMeter.LLM_TOKEN,
)
@@ -108,6 +177,7 @@ private class RequestIdProvider : GatewayProvider {
calls += 1
lastRequestId = request.requestId
lastRequestSource = request.requestSource
lastTextRequest = request as? TextProviderRequest
output.emit("""{"result":"ok"}""".encodeToByteArray())
return ProviderUsage(
meter = UsageMeter.LLM_TOKEN,
@@ -158,5 +228,9 @@ private object RequestIdUsage : GatewayUsagePort {
private val REQUEST_ID_PRINCIPAL = GatewayPrincipal(
userId = "00000000-0000-0000-0000-000000000001",
grantId = "00000000-0000-0000-0000-000000000003",
scopes = setOf(GatewayCapability.AI),
scopes = setOf(
GatewayCapability.POLISH,
GatewayCapability.AI,
GatewayCapability.AGENT,
),
)
@@ -114,15 +114,23 @@ private fun replayService(
asrProviderTimeoutMillis = timeoutMillis,
)
private fun replayRequest() = TextProviderRequest(
requestId = REPLAY_ID,
capability = GatewayCapability.AI,
input = "hello",
context = null,
maxOutputTokens = 32,
temperature = 0.2,
stream = false,
)
private fun replayRequest(): TextProviderRequest {
val executionPolicy = GatewayTaskPolicyResolver().resolve(
GatewayCapability.AI,
requestedTaskKind = null,
requestedMaxOutputTokens = 32,
)
return TextProviderRequest(
requestId = REPLAY_ID,
capability = GatewayCapability.AI,
executionPolicy = executionPolicy,
input = "hello",
context = null,
maxOutputTokens = executionPolicy.maxOutputTokens,
temperature = 0.2,
stream = false,
)
}
private class ReplayCredits(
private val failRelease: Boolean = false,
@@ -217,16 +217,24 @@ private fun service(
usageRecords = usageRecords,
)
private fun request(requestSource: GatewayRequestSource? = null) = TextProviderRequest(
requestId = "request-123",
capability = GatewayCapability.AI,
input = "hello",
context = null,
maxOutputTokens = 32,
temperature = 0.2,
stream = false,
requestSource = requestSource,
)
private fun request(requestSource: GatewayRequestSource? = null): TextProviderRequest {
val executionPolicy = GatewayTaskPolicyResolver().resolve(
GatewayCapability.AI,
requestedTaskKind = null,
requestedMaxOutputTokens = 32,
)
return TextProviderRequest(
requestId = "request-123",
capability = GatewayCapability.AI,
executionPolicy = executionPolicy,
input = "hello",
context = null,
maxOutputTokens = executionPolicy.maxOutputTokens,
temperature = 0.2,
stream = false,
requestSource = requestSource,
)
}
private class FakeCredits(
private val failSettle: Boolean = false,
@@ -0,0 +1,92 @@
package com.osglab.account.features.gateway.services
import com.osglab.account.features.gateway.models.GatewayCapability
import com.osglab.account.features.gateway.models.GatewayModelProfile
import com.osglab.account.features.gateway.models.GatewayReasoningEffort
import com.osglab.account.features.gateway.models.GatewayTaskKind
import com.osglab.account.features.gateway.models.GatewayThinkingMode
import com.osglab.account.features.gateway.models.GatewayToolsMode
import com.osglab.account.features.gateway.models.GatewayWebSearchMode
import io.kotest.assertions.throwables.shouldThrow
import io.kotest.core.spec.style.StringSpec
import io.kotest.matchers.shouldBe
class GatewayTaskPolicyResolverTest : StringSpec({
val resolver = GatewayTaskPolicyResolver()
"uses backward-compatible task defaults" {
resolver.resolve(GatewayCapability.POLISH, null, 512).taskKind shouldBe
GatewayTaskKind.DICTATION_POLISH
resolver.resolve(GatewayCapability.AI, null, 512).taskKind shouldBe
GatewayTaskKind.AI_QUESTION
resolver.resolve(GatewayCapability.AGENT, null, 512).taskKind shouldBe
GatewayTaskKind.AGENT_PLANNING
}
"forces all low-latency transform tasks to disable costly capabilities" {
val tasks = listOf(
GatewayCapability.POLISH to GatewayTaskKind.DICTATION_POLISH,
GatewayCapability.POLISH to GatewayTaskKind.TRANSLATION,
GatewayCapability.POLISH to GatewayTaskKind.EDIT_LAST_INPUT,
GatewayCapability.AI to GatewayTaskKind.CLIPBOARD_TRANSFORM,
GatewayCapability.AI to GatewayTaskKind.CUSTOM_SKILL,
)
tasks.forEach { (capability, taskKind) ->
val policy = resolver.resolve(capability, taskKind, 512)
policy.modelProfile shouldBe GatewayModelProfile.LOW_LATENCY
policy.thinking shouldBe GatewayThinkingMode.DISABLED
policy.reasoningEffort shouldBe null
policy.webSearch shouldBe GatewayWebSearchMode.DISABLED
policy.tools shouldBe GatewayToolsMode.DISABLED
policy.allowEmptyContentRetry shouldBe false
}
}
"enables explicit high-effort reasoning only for question and agent tasks" {
listOf(
GatewayCapability.AI to GatewayTaskKind.AI_QUESTION,
GatewayCapability.AGENT to GatewayTaskKind.AGENT_PLANNING,
).forEach { (capability, taskKind) ->
val policy = resolver.resolve(capability, taskKind, 512)
policy.modelProfile shouldBe GatewayModelProfile.REASONING
policy.thinking shouldBe GatewayThinkingMode.ENABLED
policy.reasoningEffort shouldBe GatewayReasoningEffort.HIGH
policy.webSearch shouldBe GatewayWebSearchMode.DISABLED
policy.tools shouldBe GatewayToolsMode.DISABLED
policy.allowEmptyContentRetry shouldBe true
}
}
"rejects capability and task mismatches without inspecting content" {
shouldThrow<IllegalArgumentException> {
resolver.resolve(
GatewayCapability.POLISH,
GatewayTaskKind.AGENT_PLANNING,
512,
)
}
shouldThrow<IllegalArgumentException> {
resolver.resolve(
GatewayCapability.AI,
GatewayTaskKind.TRANSLATION,
512,
)
}
}
"applies server output budgets" {
resolver.resolve(
GatewayCapability.POLISH,
GatewayTaskKind.DICTATION_POLISH,
4_096,
).maxOutputTokens shouldBe 512
resolver.resolve(
GatewayCapability.AI,
GatewayTaskKind.CLIPBOARD_TRANSFORM,
4_096,
).maxOutputTokens shouldBe 2_048
}
})