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
1 change: 1 addition & 0 deletions gradle/libs.versions.toml
Original file line number Diff line number Diff line change
Expand Up @@ -93,6 +93,7 @@ kotlinx-coroutines-test = { module = "org.jetbrains.kotlinx:kotlinx-coroutines-t
kotlinx-datetime = { module = "org.jetbrains.kotlinx:kotlinx-datetime", version.ref = "kotlinx-datetime" }
kotlinx-serialization = { module = "org.jetbrains.kotlinx:kotlinx-serialization-json", version.ref = "kotlinx-serialization" }
ktor-serialization-gson = { module = "io.ktor:ktor-serialization-gson", version.ref = "ktor" }
ktor-serialization-kotlinx-json = { module = "io.ktor:ktor-serialization-kotlinx-json", version.ref = "ktor" }
ktor-server-call-logging = { module = "io.ktor:ktor-server-call-logging", version.ref = "ktor" }
ktor-server-content-negotiation = { module = "io.ktor:ktor-server-content-negotiation", version.ref = "ktor" }
ktor-server-core = { module = "io.ktor:ktor-server-core", version.ref = "ktor" }
Expand Down
4 changes: 3 additions & 1 deletion webserver/build.gradle.kts
Original file line number Diff line number Diff line change
Expand Up @@ -16,6 +16,7 @@

plugins {
kotlin("jvm")
kotlin("plugin.serialization")
id("application")
id("java-library")
id("maven-publish")
Expand Down Expand Up @@ -44,13 +45,14 @@ sourceSets {
dependencies {
implementation(project(":google-adk-kotlin-core"))
implementation(libs.kotlinx.datetime)
implementation(libs.kotlinx.serialization)

implementation(libs.graphviz.java)

implementation(libs.opentelemetry.api)
implementation(libs.opentelemetry.sdk)

implementation(libs.ktor.serialization.gson)
implementation(libs.ktor.serialization.kotlinx.json)
implementation(libs.ktor.server.call.logging)
implementation(libs.ktor.server.content.negotiation)
implementation(libs.ktor.server.core)
Expand Down
Original file line number Diff line number Diff line change
Expand Up @@ -16,8 +16,10 @@

package com.google.adk.kt.webserver

import com.google.adk.kt.annotations.FrameworkInternalApi
import com.google.adk.kt.artifacts.ArtifactService
import com.google.adk.kt.runners.Runner
import com.google.adk.kt.serialization.adkJson
import com.google.adk.kt.sessions.SessionService
import com.google.adk.kt.telemetry.TelemetryConfig
import com.google.adk.kt.webserver.AdkWebServer.StatusAwareLogger
Expand All @@ -32,10 +34,7 @@ import com.google.adk.kt.webserver.routes.sessionRoutes
import com.google.adk.kt.webserver.routes.staticRoutes
import com.google.adk.kt.webserver.telemetry.ApiServerSpanExporter
import com.google.adk.kt.webserver.telemetry.OpenTelemetryConfig
import com.google.gson.TypeAdapter
import com.google.gson.stream.JsonReader
import com.google.gson.stream.JsonWriter
import io.ktor.serialization.gson.gson
import io.ktor.serialization.kotlinx.json.json
import io.ktor.server.application.Application
import io.ktor.server.application.call
import io.ktor.server.application.install
Expand All @@ -49,7 +48,6 @@ import io.ktor.server.request.uri
import io.ktor.server.response.respondText
import io.ktor.server.routing.get
import io.ktor.server.routing.routing
import kotlinx.datetime.Instant
import org.slf4j.Logger
import org.slf4j.LoggerFactory
import org.slf4j.event.Level
Expand Down Expand Up @@ -123,24 +121,6 @@ class AdkWebServer(
logger.info("Ktor server stopped")
}

class InstantTypeAdapter : TypeAdapter<Instant>() {
override fun write(out: JsonWriter, value: Instant?) {
if (value == null) {
out.nullValue()
} else {
out.value(value.toEpochMilliseconds())
}
}

override fun read(reader: JsonReader): Instant? {
if (reader.peek() == com.google.gson.stream.JsonToken.NULL) {
reader.nextNull()
return null
}
return Instant.fromEpochMilliseconds(reader.nextLong())
}
}

public class StatusAwareLogger(private val delegate: Logger) : Logger by delegate {
override fun info(msg: String?) {
if (msg != null && msg.contains("Status: 5")) {
Expand All @@ -152,6 +132,7 @@ class AdkWebServer(
}
}

@OptIn(FrameworkInternalApi::class)
fun Application.adkModule(
sessionService: SessionService,
artifactService: ArtifactService,
Expand All @@ -169,12 +150,7 @@ fun Application.adkModule(
"Status: $status, HTTP method: $httpMethod, URI: $uri"
}
}
install(ContentNegotiation) {
gson {
setPrettyPrinting()
registerTypeAdapter(Instant::class.java, AdkWebServer.InstantTypeAdapter())
}
}
install(ContentNegotiation) { json(adkJson) }

val otelConfig = OpenTelemetryConfig(apiServerSpanExporter)
val sdkTracerProvider = otelConfig.sdkTracerProvider()
Expand Down
Original file line number Diff line number Diff line change
Expand Up @@ -17,25 +17,30 @@
package com.google.adk.kt.webserver.models

import com.google.adk.kt.types.Content
import kotlinx.serialization.Contextual
import kotlinx.serialization.Serializable

@Serializable
data class AgentRunRequest(
val appName: String,
val userId: String,
val sessionId: String? = null,
val newMessage: Content? = null,
val streaming: Boolean = false,
val stateDelta: Map<String, Any>? = null,
val stateDelta: Map<String, @Contextual Any>? = null,
val invocationId: String? = null,
)

@Serializable
data class RunRequest(val agentId: String, val input: String, val sessionId: String? = null)

data class RunResponse(val output: String, val sessionId: String)
@Serializable data class RunResponse(val output: String, val sessionId: String)

data class TurnModel(val role: String, val content: String)
@Serializable data class TurnModel(val role: String, val content: String)

data class SessionModel(val sessionId: String, val turnHistory: List<TurnModel>)
@Serializable data class SessionModel(val sessionId: String, val turnHistory: List<TurnModel>)

@Serializable
data class ErrorResponse(val error: String, val message: String, val details: String? = null)

/**
Expand All @@ -45,13 +50,14 @@ data class ErrorResponse(val error: String, val message: String, val details: St
* @property content The JSON string or text content of the event.
* @property timestamp The ISO-8601 timestamp of the event.
*/
data class SseModel(val type: String, val content: String, val timestamp: String)
@Serializable data class SseModel(val type: String, val content: String, val timestamp: String)

@Serializable
data class SessionDto(
val id: String?,
val appName: String,
val userId: String,
val state: Map<String, Any>?,
val state: Map<String, @Contextual Any>?,
val events: List<com.google.adk.kt.events.Event>?,
val lastUpdateTime: Long?,
)
Original file line number Diff line number Diff line change
Expand Up @@ -18,12 +18,13 @@ package com.google.adk.kt.webserver.routes

import com.google.adk.kt.agents.RunConfig
import com.google.adk.kt.agents.StreamingMode
import com.google.adk.kt.annotations.FrameworkInternalApi
import com.google.adk.kt.artifacts.ArtifactService
import com.google.adk.kt.runners.InMemoryRunner
import com.google.adk.kt.serialization.adkJson
import com.google.adk.kt.sessions.SessionService
import com.google.adk.kt.webserver.loaders.AgentLoader
import com.google.adk.kt.webserver.models.AgentRunRequest
import com.google.gson.Gson
import io.ktor.http.ContentType
import io.ktor.http.HttpStatusCode
import io.ktor.server.application.call
Expand All @@ -37,7 +38,9 @@ import io.ktor.utils.io.writeStringUtf8
import java.util.UUID
import kotlinx.coroutines.flow.collect
import kotlinx.coroutines.flow.toList
import kotlinx.serialization.encodeToString

@OptIn(FrameworkInternalApi::class)
fun Route.runRoutes(
agentLoader: AgentLoader,
sessionService: SessionService,
Expand Down Expand Up @@ -111,7 +114,7 @@ fun Route.runRoutes(
runConfig,
)
.collect { event ->
val data = Gson().toJson(event)
val data = adkJson.encodeToString(event)
writeStringUtf8("data: $data\n\n")
flush()
}
Expand Down
Original file line number Diff line number Diff line change
Expand Up @@ -179,8 +179,8 @@ class AdkWebServerTest {
val response = client.get("/api/test-serialize")
assertThat(response.status).isEqualTo(HttpStatusCode.OK)
val body = response.bodyAsText()
assertThat(body).contains("\"output\": \"Ok output\"")
assertThat(body).contains("\"sessionId\": \"test-session\"")
assertThat(body).contains("\"output\":\"Ok output\"")
assertThat(body).contains("\"sessionId\":\"test-session\"")
}

@Test
Expand All @@ -196,8 +196,10 @@ class AdkWebServerTest {
}
assertThat(response.status).isEqualTo(HttpStatusCode.OK)
val body = response.bodyAsText()
println("RESPONSE BODY: $body")
assertThat(body).isNotEmpty()
// adkJson has encodeDefaults=false, so default-false partial/interrupted are omitted.
assertThat(body).contains("\"turnComplete\":true")
assertThat(body).doesNotContain("\"partial\"")
assertThat(body).doesNotContain("\"interrupted\"")
}

@Test
Expand All @@ -213,5 +215,11 @@ class AdkWebServerTest {
}
assertThat(response.status).isEqualTo(HttpStatusCode.OK)
assertThat(response.headers["Content-Type"]).contains("text/event-stream")
// The SSE stream serializes events with adkJson too, so the shape matches /run.
val body = response.bodyAsText()
assertThat(body).contains("data: ")
assertThat(body).contains("\"turnComplete\":true")
assertThat(body).doesNotContain("\"partial\"")
assertThat(body).doesNotContain("\"interrupted\"")
}
}
Original file line number Diff line number Diff line change
Expand Up @@ -17,12 +17,14 @@
package com.google.adk.kt.webserver.routes

import com.google.adk.kt.agents.BaseAgent
import com.google.adk.kt.annotations.FrameworkInternalApi
import com.google.adk.kt.serialization.adkJson
import com.google.adk.kt.webserver.loaders.AgentLoader
import com.google.common.truth.Truth.assertThat
import io.ktor.client.request.get
import io.ktor.client.statement.bodyAsText
import io.ktor.http.HttpStatusCode
import io.ktor.serialization.gson.gson
import io.ktor.serialization.kotlinx.json.json
import io.ktor.server.application.install
import io.ktor.server.plugins.contentnegotiation.ContentNegotiation
import io.ktor.server.routing.routing
Expand All @@ -31,6 +33,7 @@ import org.junit.Test
import org.junit.runner.RunWith
import org.junit.runners.JUnit4

@OptIn(FrameworkInternalApi::class)
@RunWith(JUnit4::class)
class AppRoutesTest {

Expand All @@ -44,7 +47,7 @@ class AppRoutesTest {
fun listApps_returnsAppList() = testApplication {
val fakeLoader = FakeAgentLoader(agentList = listOf("app1", "app2"))
application {
install(ContentNegotiation) { gson { setPrettyPrinting() } }
install(ContentNegotiation) { json(adkJson) }
routing { appRoutes(fakeLoader) }
}

Expand All @@ -60,7 +63,7 @@ class AppRoutesTest {
fun listApps_empty_returnsEmptyList() = testApplication {
val fakeLoader = FakeAgentLoader(agentList = emptyList())
application {
install(ContentNegotiation) { gson { setPrettyPrinting() } }
install(ContentNegotiation) { json(adkJson) }
routing { appRoutes(fakeLoader) }
}

Expand Down
Loading
Loading