diff --git a/build.sbt b/build.sbt index c9fdeb2..27abad2 100644 --- a/build.sbt +++ b/build.sbt @@ -4,7 +4,7 @@ inThisBuild( Seq( crossScalaVersions := Seq(scala213Version, scala3Version), scalaVersion := scala213Version, - tlBaseVersion := "0.3", + tlBaseVersion := "0.4", organizationName := "Christopher Davenport", startYear := Some(2023), licenses := Seq(License.MIT), @@ -93,7 +93,7 @@ lazy val codeGeneratorSbt2 = project libraryDependencies ++= Seq( "com.thesamet.scalapb" %% "compilerplugin" % scalapbSbt2Version ), - tlVersionIntroduced := Map("3" -> "0.3.1"), + tlVersionIntroduced := Map("3" -> "0.4.0"), ) .disablePlugins(ScalafixPlugin) @@ -115,7 +115,7 @@ lazy val codeGeneratorPlugin = project case "2.12" => Some(8) case _ => Some(17) }), - tlVersionIntroduced := Map("3" -> "0.3.1"), + tlVersionIntroduced := Map("3" -> "0.4.0"), tlFatalWarnings := false, buildInfoPackage := "org.http4s.grpc.sbt", buildInfoOptions += BuildInfoOption.PackagePrivate, diff --git a/codegen/generator/src/main/scala/org/http4s/grpc/generator/Http4sGrpcServicePrinter.scala b/codegen/generator/src/main/scala/org/http4s/grpc/generator/Http4sGrpcServicePrinter.scala index 7c09090..9f127c3 100644 --- a/codegen/generator/src/main/scala/org/http4s/grpc/generator/Http4sGrpcServicePrinter.scala +++ b/codegen/generator/src/main/scala/org/http4s/grpc/generator/Http4sGrpcServicePrinter.scala @@ -67,11 +67,11 @@ class Http4sGrpcServicePrinter(service: ServiceDescriptor, di: DescriptorImplici fp.add(lines: _*) } - private[this] def serviceMethodSignature(method: MethodDescriptor) = { + private[this] def serviceMethodSignature(method: MethodDescriptor, ctxType: String) = { val scalaInType = scalaType(method.getInputType, method.inputType) val scalaOutType = scalaType(method.getOutputType, method.outputType) - val ctx = s"ctx: $Ctx" + val ctx = s"ctx: $ctxType" s"def ${method.name}" + (method.streamType match { case StreamType.Unary => s"(request: $scalaInType, $ctx): F[$scalaOutType]" @@ -91,45 +91,61 @@ class Http4sGrpcServicePrinter(service: ServiceDescriptor, di: DescriptorImplici case StreamType.Bidirectional => "streamToStream" } - private[this] def createClientCall(method: MethodDescriptor) = { + private[this] def createClientCall(method: MethodDescriptor, headers: String) = { val encode = codec(method.getInputType, method.inputType) val decode = codec(method.getOutputType, method.outputType) val serviceName = method.getService.getFullName val methodName = method.getName s"""$ClientGrpc.${handleMethod( method - )}($encode, $decode, "$serviceName", "$methodName", maxMessageSize)(client, baseUri)(request, ctx)""" + )}($encode, $decode, "$serviceName", "$methodName", maxMessageSize)(client, baseUri)(request, $headers)""" } private[this] def serviceMethodImplementation(method: MethodDescriptor): PrinterEndo = { p => - p.add(serviceMethodSignature(method) + " = {") + p.add(serviceMethodSignature(method, Headers) + " = {") .indent - .add(s"${createClientCall(method)}") + .add(s"${createClientCall(method, "ctx")}") .outdent .add("}") } - private[this] def serviceBindingImplementation(method: MethodDescriptor): PrinterEndo = { p => - // val serviceCall = s"serviceImpl.${method.name}" - // val eval = if (method.isServerStreaming) s"$Stream.eval(mkCtx(m))" else "mkCtx(m)" + private[this] def contextServiceMethodImplementation(method: MethodDescriptor): PrinterEndo = { + p => + val call = createClientCall(method, "_") + val impl = + if (method.isServerStreaming) s"$Stream.eval(mkHeaders(ctx)).flatMap($call)" + else s"mkHeaders(ctx).flatMap($call)" + + p.add(serviceMethodSignature(method, "A") + " = {") + .indent + .add(impl) + .outdent + .add("}") + } + private[this] def serviceBindingImplementation(method: MethodDescriptor): PrinterEndo = { p => val decode = codec(method.getInputType, method.inputType) val encode = codec(method.getOutputType, method.outputType) val serviceName = method.getService.getFullName val methodName = method.getName - p.add(s""".combineK($ServerGrpc.${handleMethod( - method - )}($decode, $encode, "$serviceName", "$methodName", maxMessageSize)(serviceImpl.${method.name}(_, _)))""") + p.add( + s""".combineK($ServerGrpc.${handleMethod( + method + )}($decode, $encode, "$serviceName", "$methodName", maxMessageSize, mkCtx)(serviceImpl.${method.name}(_, _)))""" + ) } private[this] def serviceMethods: PrinterEndo = _.call(service.methods.map { method => - generateScalaDoc(method).andThen(_.add(serviceMethodSignature(method)).newline) + generateScalaDoc(method).andThen(_.add(serviceMethodSignature(method, "A")).newline) }: _*) private[this] def serviceMethodImplementations: PrinterEndo = _.call(service.methods.map(serviceMethodImplementation): _*) + private[this] def contextServiceMethodImplementations: PrinterEndo = + _.call(service.methods.map(contextServiceMethodImplementation): _*) + private[this] def serviceBindingImplementations: PrinterEndo = _.add(s"$ServerGrpc.precondition[F]").indent .call(service.methods.map(serviceBindingImplementation): _*) @@ -143,7 +159,11 @@ class Http4sGrpcServicePrinter(service: ServiceDescriptor, di: DescriptorImplici private[this] def serviceTrait: PrinterEndo = _.call(generateScalaDoc(service)) - .add(s"trait $serviceName[F[_]] {") + .add(s"trait $serviceName[F[_]] extends $serviceType.WithContext[F, $Headers]") + + private[this] def contextServiceTrait: PrinterEndo = + _.call(generateScalaDoc(service)) + .add(s"trait WithContext[F[_], A] {") .newline .indent .call(serviceMethods) @@ -152,8 +172,12 @@ class Http4sGrpcServicePrinter(service: ServiceDescriptor, di: DescriptorImplici private[this] def serviceObject: PrinterEndo = _.add(s"object $serviceName {").indent.newline + .call(contextServiceTrait) + .newline .call(serviceClient) .newline + .call(contextServiceClient) + .newline .call(serviceBinding) .outdent .newline @@ -171,12 +195,32 @@ class Http4sGrpcServicePrinter(service: ServiceDescriptor, di: DescriptorImplici .outdent .add("}") + private[this] def contextServiceClient: PrinterEndo = + _.add( + s"def fromClient[F[_]: $Concurrent, A](client: $Client[F], baseUri: $Uri, mkHeaders: A => F[$Headers]): $serviceType.WithContext[F, A] = fromClient(client, baseUri, mkHeaders, $DefaultMaxMessageSize)" + ).newline + .add( + s"def fromClient[F[_]: $Concurrent, A](client: $Client[F], baseUri: $Uri, mkHeaders: A => F[$Headers], maxMessageSize: Int): $serviceType.WithContext[F, A] = new $serviceType.WithContext[F, A] {" + ) + .indent + .call(contextServiceMethodImplementations) + .outdent + .add("}") + private[this] def serviceBinding: PrinterEndo = _.add( s"def toRoutes[F[_]: $Temporal](serviceImpl: $serviceType[F]): $HttpRoutes[F] = toRoutes(serviceImpl, $DefaultMaxMessageSize)" ).newline .add( - s"def toRoutes[F[_]: $Temporal](serviceImpl: $serviceType[F], maxMessageSize: Int): $HttpRoutes[F] = {" + s"def toRoutes[F[_]: $Temporal](serviceImpl: $serviceType[F], maxMessageSize: Int): $HttpRoutes[F] = toRoutes[F, $Headers](serviceImpl, (request: $Request[F]) => request.headers.pure[F], maxMessageSize)" + ) + .newline + .add( + s"def toRoutes[F[_]: $Temporal, A](serviceImpl: $serviceType.WithContext[F, A], mkCtx: $Request[F] => F[A]): $HttpRoutes[F] = toRoutes(serviceImpl, mkCtx, $DefaultMaxMessageSize)" + ) + .newline + .add( + s"def toRoutes[F[_]: $Temporal, A](serviceImpl: $serviceType.WithContext[F, A], mkCtx: $Request[F] => F[A], maxMessageSize: Int): $HttpRoutes[F] = {" ) .indent .call(serviceBindingImplementations) @@ -208,7 +252,8 @@ object Http4sGrpcServicePrinter { // / - val Ctx = s"$http4sPkg.Headers" + val Headers = s"$http4sPkg.Headers" + val Request = s"$http4sPkg.Request" val Concurrent = s"$effPkg.Concurrent" val Temporal = s"$effPkg.Temporal" diff --git a/codegen/testing/src/test/scala/hello/world/WithContextSuite.scala b/codegen/testing/src/test/scala/hello/world/WithContextSuite.scala new file mode 100644 index 0000000..36880ae --- /dev/null +++ b/codegen/testing/src/test/scala/hello/world/WithContextSuite.scala @@ -0,0 +1,139 @@ +/* + * Copyright (c) 2023 Christopher Davenport + * + * Permission is hereby granted, free of charge, to any person obtaining a copy of + * this software and associated documentation files (the "Software"), to deal in + * the Software without restriction, including without limitation the rights to + * use, copy, modify, merge, publish, distribute, sublicense, and/or sell copies of + * the Software, and to permit persons to whom the Software is furnished to do so, + * subject to the following conditions: + * + * The above copyright notice and this permission notice shall be included in all + * copies or substantial portions of the Software. + * + * THE SOFTWARE IS PROVIDED "AS IS", WITHOUT WARRANTY OF ANY KIND, EXPRESS OR + * IMPLIED, INCLUDING BUT NOT LIMITED TO THE WARRANTIES OF MERCHANTABILITY, FITNESS + * FOR A PARTICULAR PURPOSE AND NONINFRINGEMENT. IN NO EVENT SHALL THE AUTHORS OR + * COPYRIGHT HOLDERS BE LIABLE FOR ANY CLAIM, DAMAGES OR OTHER LIABILITY, WHETHER + * IN AN ACTION OF CONTRACT, TORT OR OTHERWISE, ARISING FROM, OUT OF OR IN + * CONNECTION WITH THE SOFTWARE OR THE USE OR OTHER DEALINGS IN THE SOFTWARE. + */ + +package hello.world + +import cats.effect.IO +import cats.syntax.all._ +import fs2.Stream +import munit._ +import org.http4s._ +import org.http4s.client.Client +import org.http4s.grpc.GrpcStatus._ +import org.http4s.grpc.GrpcStatusException +import org.typelevel.ci._ +import org.typelevel.vault.Key + +class WithContextSuite extends CatsEffectSuite { + val impl: TestService.WithContext[IO, String] = new TestService.WithContext[IO, String] { + def noStreaming(request: TestMessage, ctx: String): IO[TestMessage] = + IO(request.copy(a = ctx)) + + def clientStreaming(request: Stream[IO, TestMessage], ctx: String): IO[TestMessage] = + request.compile.lastOrError.map(_.copy(a = ctx)) + + def serverStreaming(request: TestMessage, ctx: String): Stream[IO, TestMessage] = + Stream.emit(request.copy(a = ctx)) + + def bothStreaming(request: Stream[IO, TestMessage], ctx: String): Stream[IO, TestMessage] = + request.map(_.copy(a = ctx)) + + def `export`(request: TestMessage, ctx: String): IO[TestMessage] = + IO(request.copy(a = ctx)) + } + + val msg: TestMessage = TestMessage("", 1, None) + + private def callAll(client: TestService.WithContext[IO, Headers]): IO[List[String]] = + List( + client.noStreaming(msg, Headers.empty), + client.clientStreaming(Stream.emit(msg), Headers.empty), + client.serverStreaming(msg, Headers.empty).compile.lastOrError, + client.bothStreaming(Stream.emit(msg), Headers.empty).compile.lastOrError, + ).traverse(_.map(_.a)) + + test("Server context is derived from request attributes") { + Key.newKey[IO, String].flatMap { userId => + val routes = TestService.toRoutes[IO, String]( + impl, + req => IO.fromOption(req.attributes.lookup(userId))(new NoSuchElementException("userId")), + ) + val withAuth = HttpRoutes[IO](req => routes.run(req.withAttribute(userId, "user-1"))) + val client = TestService.fromClient[IO](Client.fromHttpApp(withAuth.orNotFound), Uri()) + + callAll(client).assertEquals(List.fill(4)("user-1")) + } + } + + test("Server context failure is returned as the grpc status") { + val routes = TestService.toRoutes[IO, String]( + impl, + _ => IO.raiseError(Unauthenticated.withMessage("no user").toException), + ) + val client = TestService.fromClient[IO](Client.fromHttpApp(routes.orNotFound), Uri()) + + List( + client.noStreaming(msg, Headers.empty), + client.clientStreaming(Stream.emit(msg), Headers.empty), + client.serverStreaming(msg, Headers.empty).compile.lastOrError, + client.bothStreaming(Stream.emit(msg), Headers.empty).compile.lastOrError, + ).traverse(_.attemptNarrow[GrpcStatusException].map(_.leftMap(_.status))) + .assertEquals(List.fill(4)(Either.left(Unauthenticated.withMessage("no user")))) + } + + test("Client context is converted to request headers") { + val routes = TestService.toRoutes[IO, String]( + impl, + req => + IO.fromOption(req.headers.get(ci"x-user-id").map(_.head.value))( + new NoSuchElementException("x-user-id") + ), + ) + val client = TestService.fromClient[IO, String]( + Client.fromHttpApp(routes.orNotFound), + Uri(), + (userId: String) => IO.pure(Headers("x-user-id" -> userId)), + ) + + List( + client.noStreaming(msg, "user-2"), + client.clientStreaming(Stream.emit(msg), "user-2"), + client.serverStreaming(msg, "user-2").compile.lastOrError, + client.bothStreaming(Stream.emit(msg), "user-2").compile.lastOrError, + ).traverse(_.map(_.a)) + .assertEquals(List.fill(4)("user-2")) + } + + test("Headers-based impl can be served with a derived context") { + val headersImpl: TestService[IO] = new TestService[IO] { + def noStreaming(request: TestMessage, ctx: Headers): IO[TestMessage] = + IO(request.copy(a = ctx.get(ci"x-user-id").fold("")(_.head.value))) + + def clientStreaming(request: Stream[IO, TestMessage], ctx: Headers): IO[TestMessage] = + request.compile.lastOrError + + def serverStreaming(request: TestMessage, ctx: Headers): Stream[IO, TestMessage] = + Stream.emit(request) + + def bothStreaming(request: Stream[IO, TestMessage], ctx: Headers): Stream[IO, TestMessage] = + request + + def `export`(request: TestMessage, ctx: Headers): IO[TestMessage] = IO(request) + } + val routes = TestService.toRoutes[IO, Headers]( + headersImpl, + req => IO.pure(req.headers.put("x-user-id" -> "user-3")), + ) + val client = TestService.fromClient[IO](Client.fromHttpApp(routes.orNotFound), Uri()) + + client.noStreaming(msg, Headers.empty).map(_.a).assertEquals("user-3") + } +} diff --git a/core/src/main/scala/org/http4s/grpc/ServerGrpc.scala b/core/src/main/scala/org/http4s/grpc/ServerGrpc.scala index 22560d0..228a387 100644 --- a/core/src/main/scala/org/http4s/grpc/ServerGrpc.scala +++ b/core/src/main/scala/org/http4s/grpc/ServerGrpc.scala @@ -21,6 +21,7 @@ package org.http4s.grpc +import cats.Applicative import cats.Functor import cats.Monad import cats.effect._ @@ -67,6 +68,18 @@ object ServerGrpc { maxMessageSize: Int, )( // Stuff we apply at invocation f: (A, Headers) => F[B] + ): HttpRoutes[F] = + unaryToUnary(decode, encode, serviceName, methodName, maxMessageSize, headers[F])(f) + + def unaryToUnary[F[_]: Temporal, A, B, C]( // Stuff We can provide via codegen\ + decode: Decoder[A], + encode: Encoder[B], + serviceName: String, + methodName: String, + maxMessageSize: Int, + mkCtx: Request[F] => F[C], + )( // Stuff we apply at invocation + f: (A, C) => F[B] ): HttpRoutes[F] = HttpRoutes.of[F] { case req @ POST -> Root / sN / mN if sN === serviceName && mN === methodName => for { @@ -76,7 +89,7 @@ object ServerGrpc { } yield { val body = Stream .eval(codecs.Messages.decodeSingle(decode, maxMessageSize)(req.body)) - .evalMap(f(_, req.headers)) + .evalMap(a => mkCtx(req).flatMap(f(a, _))) .flatMap(codecs.Messages.encodeSingle(encode)(_)) .through(timeoutStream(_)(timeout.map(_.duration))) .onFinalizeCaseWeak(updateStatus(status)) @@ -110,6 +123,18 @@ object ServerGrpc { maxMessageSize: Int, )( // Stuff we apply at invocation f: (A, Headers) => Stream[F, B] + ): HttpRoutes[F] = + unaryToStream(decode, encode, serviceName, methodName, maxMessageSize, headers[F])(f) + + def unaryToStream[F[_]: Temporal, A, B, C]( // Stuff We can provide via codegen\ + decode: Decoder[A], + encode: Encoder[B], + serviceName: String, + methodName: String, + maxMessageSize: Int, + mkCtx: Request[F] => F[C], + )( // Stuff we apply at invocation + f: (A, C) => Stream[F, B] ): HttpRoutes[F] = HttpRoutes.of[F] { case req @ POST -> Root / sN / mN if sN === serviceName && mN === methodName => for { @@ -119,7 +144,7 @@ object ServerGrpc { } yield { val body = Stream .eval(codecs.Messages.decodeSingle(decode, maxMessageSize)(req.body)) - .flatMap(f(_, req.headers)) + .flatMap(a => Stream.eval(mkCtx(req)).flatMap(f(a, _))) .through(codecs.Messages.encode(encode)) .through(timeoutStream(_)(timeout.map(_.duration))) .onFinalizeCaseWeak(updateStatus(status)) @@ -152,6 +177,18 @@ object ServerGrpc { maxMessageSize: Int, )( // Stuff we apply at invocation f: (Stream[F, A], Headers) => F[B] + ): HttpRoutes[F] = + streamToUnary(decode, encode, serviceName, methodName, maxMessageSize, headers[F])(f) + + def streamToUnary[F[_]: Temporal, A, B, C]( // Stuff We can provide via codegen\ + decode: Decoder[A], + encode: Encoder[B], + serviceName: String, + methodName: String, + maxMessageSize: Int, + mkCtx: Request[F] => F[C], + )( // Stuff we apply at invocation + f: (Stream[F, A], C) => F[B] ): HttpRoutes[F] = HttpRoutes.of[F] { case req @ POST -> Root / sN / mN if sN === serviceName && mN === methodName => for { @@ -161,7 +198,7 @@ object ServerGrpc { } yield { val body = Stream - .eval(f(codecs.Messages.decode(decode, maxMessageSize)(req.body), req.headers)) + .eval(mkCtx(req).flatMap(f(codecs.Messages.decode(decode, maxMessageSize)(req.body), _))) .flatMap(codecs.Messages.encodeSingle(encode)(_)) .through(timeoutStream(_)(timeout.map(_.duration))) .onFinalizeCaseWeak(updateStatus(status)) @@ -197,6 +234,18 @@ object ServerGrpc { maxMessageSize: Int, )( // Stuff we apply at invocation f: (Stream[F, A], Headers) => Stream[F, B] + ): HttpRoutes[F] = + streamToStream(decode, encode, serviceName, methodName, maxMessageSize, headers[F])(f) + + def streamToStream[F[_]: Temporal, A, B, C]( // Stuff We can provide via codegen\ + decode: Decoder[A], + encode: Encoder[B], + serviceName: String, + methodName: String, + maxMessageSize: Int, + mkCtx: Request[F] => F[C], + )( // Stuff we apply at invocation + f: (Stream[F, A], C) => Stream[F, B] ): HttpRoutes[F] = HttpRoutes.of[F] { case req @ POST -> Root / sN / mN if sN === serviceName && mN === methodName => for { @@ -205,7 +254,9 @@ object ServerGrpc { timeout = req.headers.get[NamedHeaders.GrpcTimeout] } yield { - val body = f(codecs.Messages.decode(decode, maxMessageSize)(req.body), req.headers) + val body = Stream + .eval(mkCtx(req)) + .flatMap(f(codecs.Messages.decode(decode, maxMessageSize)(req.body), _)) .through(codecs.Messages.encode(encode)) .through(timeoutStream(_)(timeout.map(_.duration))) .onFinalizeCaseWeak(updateStatus(status)) @@ -265,6 +316,9 @@ object ServerGrpc { .pure[F] } + private def headers[F[_]: Applicative]: Request[F] => F[Headers] = + _.headers.pure[F] + private def timeoutStream[F[_]: Temporal, A]( s: Stream[F, A] )(timeout: Option[FiniteDuration]): Stream[F, A] =