Skip to content

Commit fa0adb3

Browse files
committed
Allow interceptor to change the effect type
1 parent d58c618 commit fa0adb3

6 files changed

Lines changed: 154 additions & 43 deletions

File tree

grpc-fs2/src/main/scala/proteus/server/Fs2ServerBackend.scala

Lines changed: 25 additions & 15 deletions
Original file line numberDiff line numberDiff line change
@@ -10,32 +10,42 @@ import io.grpc.*
1010

1111
import proteus.server.ServerInterceptor
1212

13-
class Fs2ServerBackend[F[_]: Async, Context](
14-
interceptor: ServerInterceptor[F, Stream[F, *], RequestResponseMetadata, Context],
13+
class Fs2ServerBackend[F[_]: Async, G[_], Context](
14+
interceptor: ServerInterceptor[F, G, Stream[F, *], Stream[G, *], RequestResponseMetadata, Context],
1515
dispatcher: Dispatcher[F],
1616
serverOptions: ServerOptions = ServerOptions.default
17-
) extends ServerBackend[F, Stream[F, *], Context] {
17+
) extends ServerBackend[G, Stream[G, *], Context] {
1818
def handler[Request, Response](
19-
rpc: ServerRpc[F, Stream[F, *], Context, Request, Response]
19+
rpc: ServerRpc[G, Stream[G, *], Context, Request, Response]
2020
): ServerCallHandler[Request, Response] =
2121
rpc match {
22-
case ServerRpc.Unary(_, logic) =>
22+
case ServerRpc.Unary(rpc, logic) =>
2323
Fs2ServerCallHandler[F](dispatcher, serverOptions).unaryToUnaryCallTrailers { (req, context) =>
2424
val responseMetadata = new Metadata()
25-
interceptor.unary(ctx => logic(req, ctx))(RequestResponseMetadata(context, responseMetadata)).map((_, responseMetadata))
25+
interceptor
26+
.unary(req, ctx => logic(req, ctx))(using rpc.requestCodec, rpc.responseCodec)(RequestResponseMetadata(context, responseMetadata))
27+
.map((_, responseMetadata))
2628
}
27-
case ServerRpc.ClientStreaming(_, logic) =>
29+
case ServerRpc.ClientStreaming(rpc, logic) =>
2830
Fs2ServerCallHandler[F](dispatcher, serverOptions).streamingToUnaryCallTrailers { (req, context) =>
2931
val responseMetadata = new Metadata()
30-
interceptor.unary(ctx => logic(req, ctx))(RequestResponseMetadata(context, responseMetadata)).map((_, responseMetadata))
32+
interceptor
33+
.clientStreaming[Request, Response](req => ctx => logic(req, ctx))(using rpc.requestCodec, rpc.responseCodec)(req)(
34+
RequestResponseMetadata(context, responseMetadata)
35+
)
36+
.map((_, responseMetadata))
3137
}
32-
case ServerRpc.ServerStreaming(_, logic) =>
38+
case ServerRpc.ServerStreaming(rpc, logic) =>
3339
Fs2ServerCallHandler[F](dispatcher, serverOptions).unaryToStreamingCall { (req, context) =>
34-
interceptor.stream(ctx => logic(req, ctx))(RequestResponseMetadata(context, new Metadata()))
40+
interceptor.serverStreaming(req, ctx => logic(req, ctx))(using rpc.requestCodec, rpc.responseCodec)(
41+
RequestResponseMetadata(context, new Metadata())
42+
)
3543
}
36-
case ServerRpc.BidiStreaming(_, logic) =>
44+
case ServerRpc.BidiStreaming(rpc, logic) =>
3745
Fs2ServerCallHandler[F](dispatcher, serverOptions).streamingToStreamingCall { (req, context) =>
38-
interceptor.stream(ctx => logic(req, ctx))(RequestResponseMetadata(context, new Metadata()))
46+
interceptor.bidiStreaming[Request, Response](req => ctx => logic(req, ctx))(using rpc.requestCodec, rpc.responseCodec)(req)(
47+
RequestResponseMetadata(context, new Metadata())
48+
)
3949
}
4050
}
4151
}
@@ -44,13 +54,13 @@ object Fs2ServerBackend {
4454
def apply[F[_]: Async](
4555
dispatcher: Dispatcher[F],
4656
serverOptions: ServerOptions = ServerOptions.default
47-
): Fs2ServerBackend[F, RequestResponseMetadata] =
57+
): Fs2ServerBackend[F, F, RequestResponseMetadata] =
4858
apply(ServerInterceptor.empty, dispatcher, serverOptions)
4959

5060
def apply[F[_]: Async, Context](
51-
interceptor: ServerInterceptor[F, Stream[F, *], RequestResponseMetadata, Context],
61+
interceptor: ServerContextInterceptor[F, Stream[F, *], RequestResponseMetadata, Context],
5262
dispatcher: Dispatcher[F],
5363
serverOptions: ServerOptions
54-
): Fs2ServerBackend[F, Context] =
64+
): Fs2ServerBackend[F, F, Context] =
5565
new Fs2ServerBackend(interceptor, dispatcher, serverOptions)
5666
}

grpc-zio/src/main/scala/proteus/server/ZioServerBackend.scala

Lines changed: 26 additions & 12 deletions
Original file line numberDiff line numberDiff line change
@@ -9,22 +9,36 @@ import zio.stream.*
99

1010
import proteus.server.ServerInterceptor
1111

12-
class ZioServerBackend[Context](
13-
interceptor: ServerInterceptor[IO[StatusException, *], ZStream[Any, StatusException, *], RequestContext, Context],
12+
class ZioServerBackend[R, E, Context](
13+
interceptor: ServerInterceptor[IO[StatusException, *], ZIO[R, E, *], ZStream[Any, StatusException, *], ZStream[R, E, *], RequestContext, Context],
1414
runtime: Runtime[Any] = Runtime.default
15-
) extends ServerBackend[IO[StatusException, *], ZStream[Any, StatusException, *], Context] {
15+
) extends ServerBackend[ZIO[R, E, *], ZStream[R, E, *], Context] {
1616
def handler[Request, Response](
17-
rpc: ServerRpc[IO[StatusException, *], ZStream[Any, StatusException, *], Context, Request, Response]
17+
rpc: ServerRpc[ZIO[R, E, *], ZStream[R, E, *], Context, Request, Response]
1818
): ServerCallHandler[Request, Response] =
1919
rpc match {
20-
case ServerRpc.Unary(_, logic) =>
21-
ZServerCallHandler.unaryCallHandler(runtime, (req, context) => interceptor.unary(ctx => logic(req, ctx))(context))
22-
case ServerRpc.ClientStreaming(_, logic) =>
23-
ZServerCallHandler.clientStreamingCallHandler(runtime, (req, context) => interceptor.unary(ctx => logic(req, ctx))(context))
24-
case ServerRpc.ServerStreaming(_, logic) =>
25-
ZServerCallHandler.serverStreamingCallHandler(runtime, (req, context) => interceptor.stream(ctx => logic(req, ctx))(context))
26-
case ServerRpc.BidiStreaming(_, logic) =>
27-
ZServerCallHandler.bidiCallHandler(runtime, (req, context) => interceptor.stream(ctx => logic(req, ctx))(context))
20+
case ServerRpc.Unary(rpc, logic) =>
21+
ZServerCallHandler.unaryCallHandler(
22+
runtime,
23+
(req, context) => interceptor.unary(req, ctx => logic(req, ctx))(using rpc.requestCodec, rpc.responseCodec)(context)
24+
)
25+
case ServerRpc.ClientStreaming(rpc, logic) =>
26+
ZServerCallHandler.clientStreamingCallHandler(
27+
runtime,
28+
(req, context) =>
29+
interceptor.clientStreaming[Request, Response](req => ctx => logic(req, ctx))(using rpc.requestCodec, rpc.responseCodec)(req)(context)
30+
)
31+
case ServerRpc.ServerStreaming(rpc, logic) =>
32+
ZServerCallHandler.serverStreamingCallHandler(
33+
runtime,
34+
(req, context) => interceptor.serverStreaming(req, ctx => logic(req, ctx))(using rpc.requestCodec, rpc.responseCodec)(context)
35+
)
36+
case ServerRpc.BidiStreaming(rpc, logic) =>
37+
ZServerCallHandler.bidiCallHandler(
38+
runtime,
39+
(req, context) =>
40+
interceptor.bidiStreaming[Request, Response](req => ctx => logic(req, ctx))(using rpc.requestCodec, rpc.responseCodec)(req)(context)
41+
)
2842
}
2943
}
3044

grpc-zio/src/test/scala/proteus/ZioBackendSpec.scala

Lines changed: 57 additions & 3 deletions
Original file line numberDiff line numberDiff line change
@@ -2,8 +2,7 @@ package proteus
22

33
import java.util.concurrent.TimeUnit
44

5-
import io.grpc.Metadata
6-
import io.grpc.StatusException
5+
import io.grpc.{ServerInterceptor as _, *}
76
import io.grpc.netty.{NettyChannelBuilder, NettyServerBuilder}
87
import io.grpc.protobuf.services.ProtoReflectionServiceV1
98
import scalapb.zio_grpc.{RequestContext, ZChannel}
@@ -13,7 +12,7 @@ import zio.test.*
1312

1413
import proteus.GrpcTestUtils.*
1514
import proteus.client.ZioClientBackend
16-
import proteus.server.{ServerService, ZioServerBackend}
15+
import proteus.server.*
1716

1817
object ZioBackendSpec extends ZIOSpecDefault {
1918

@@ -217,6 +216,61 @@ object ZioBackendSpec extends ZIOSpecDefault {
217216
server.shutdown().awaitTermination(5, TimeUnit.SECONDS)
218217
channel.shutdown().awaitTermination(5, TimeUnit.SECONDS)
219218
}.ignore)
219+
},
220+
test("should handle server interceptor that changes the effect type") {
221+
val backend = ZioServerBackend(
222+
new ServerInterceptor[
223+
IO[StatusException, *],
224+
IO[String, *],
225+
ZStream[Any, StatusException, *],
226+
ZStream[Any, String, *],
227+
RequestContext,
228+
RequestContext
229+
] {
230+
def unary[Req: ProtobufCodec, Resp: ProtobufCodec](
231+
request: Req,
232+
io: RequestContext => IO[String, Resp]
233+
): (RequestContext => IO[StatusException, Resp]) =
234+
ctx => io(ctx).mapError(error => Status.INTERNAL.withDescription(error).asException())
235+
def clientStreaming[Req: ProtobufCodec, Resp: ProtobufCodec](
236+
io: ZStream[Any, String, Req] => RequestContext => IO[String, Resp]
237+
): (ZStream[Any, StatusException, Req] => RequestContext => IO[StatusException, Resp]) =
238+
stream => ctx => io(stream.mapError(_.getMessage))(ctx).mapError(error => Status.INTERNAL.withDescription(error).asException())
239+
def serverStreaming[Req: ProtobufCodec, Resp: ProtobufCodec](
240+
request: Req,
241+
io: RequestContext => ZStream[Any, String, Resp]
242+
): (RequestContext => ZStream[Any, StatusException, Resp]) =
243+
ctx => io(ctx).mapError(error => Status.INTERNAL.withDescription(error).asException())
244+
def bidiStreaming[Req: ProtobufCodec, Resp: ProtobufCodec](
245+
io: ZStream[Any, String, Req] => RequestContext => ZStream[Any, String, Resp]
246+
): (ZStream[Any, StatusException, Req] => RequestContext => ZStream[Any, StatusException, Resp]) =
247+
stream => ctx => io(stream.mapError(_.getMessage))(ctx).mapError(error => Status.INTERNAL.withDescription(error).asException())
248+
},
249+
Runtime.default
250+
)
251+
val serverService = ServerService(using backend)
252+
.rpc(complexRpc, _ => ZIO.fail("boom"))
253+
.build(testService)
254+
255+
val port = 7011
256+
val server = NettyServerBuilder.forPort(port).addService(serverService).build().start()
257+
val channel = NettyChannelBuilder.forAddress("localhost", port).usePlaintext().build()
258+
val zChannel = ZChannel(channel, Seq.empty)
259+
val clientBackend = new ZioClientBackend(zChannel)
260+
261+
val program = for {
262+
client <- clientBackend.client(complexRpc, testService)
263+
264+
response1 <- client(sampleRequest).either
265+
} yield response1
266+
267+
program
268+
.flatMap(result => assertTrue(result.left.map(_.getStatus().getDescription()) == Left("boom")))
269+
.ensuring(ZIO.attempt {
270+
server.shutdown().awaitTermination(5, TimeUnit.SECONDS)
271+
channel.shutdown().awaitTermination(5, TimeUnit.SECONDS)
272+
}.ignore)
273+
.provide()
220274
}
221275
)
222276
}

grpc/src/main/scala/proteus/server/DirectServerBackend.scala

Lines changed: 7 additions & 4 deletions
Original file line numberDiff line numberDiff line change
@@ -3,19 +3,22 @@ package server
33

44
import io.grpc.{Metadata, ServerCall, ServerCallHandler, Status}
55

6-
class DirectServerBackend[Context](interceptor: ServerInterceptor[[A] =>> A, [A] =>> A, RequestResponseMetadata, Context])
6+
class DirectServerBackend[Context](interceptor: ServerContextInterceptor[[A] =>> A, [A] =>> A, RequestResponseMetadata, Context])
77
extends ServerBackend[[A] =>> A, [A] =>> A, Context] {
88
def handler[Request, Response](rpc: ServerRpc[[A] =>> A, [A] =>> A, Context, Request, Response]): ServerCallHandler[Request, Response] =
99
rpc match {
10-
case server.ServerRpc.Unary(_, logic) =>
10+
case server.ServerRpc.Unary(rpc, logic) =>
1111
new ServerCallHandler[Request, Response] {
1212
def startCall(call: ServerCall[Request, Response], headers: Metadata): ServerCall.Listener[Request] = {
1313
call.request(1)
1414
new ServerCall.Listener[Request] {
1515
override def onMessage(message: Request): Unit =
1616
try {
1717
val responseMetadata = new Metadata()
18-
val response = interceptor.unary(ctx => logic(message, ctx))(RequestResponseMetadata(headers, responseMetadata))
18+
val response =
19+
interceptor.unary(message, ctx => logic(message, ctx))(using rpc.requestCodec, rpc.responseCodec)(
20+
RequestResponseMetadata(headers, responseMetadata)
21+
)
1922
call.sendHeaders(new Metadata())
2023
call.sendMessage(response)
2124
call.close(Status.OK, responseMetadata)
@@ -26,7 +29,7 @@ class DirectServerBackend[Context](interceptor: ServerInterceptor[[A] =>> A, [A]
2629
}
2730
}
2831
}
29-
case _ =>
32+
case _ =>
3033
throw new UnsupportedOperationException("The direct backend only supports unary RPCs")
3134
}
3235
}

grpc/src/main/scala/proteus/server/FutureServerBackend.scala

Lines changed: 5 additions & 2 deletions
Original file line numberDiff line numberDiff line change
@@ -5,7 +5,7 @@ import scala.concurrent.Future
55

66
import io.grpc.{Metadata, ServerCall, ServerCallHandler, Status}
77

8-
class FutureServerBackend[Context](interceptor: ServerInterceptor[Future, Future, RequestResponseMetadata, Context])
8+
class FutureServerBackend[Context](interceptor: ServerContextInterceptor[Future, Future, RequestResponseMetadata, Context])
99
extends ServerBackend[Future, Future, Context] {
1010
def handler[Request, Response](
1111
rpc: ServerRpc[Future, Future, Context, Request, Response]
@@ -19,7 +19,10 @@ class FutureServerBackend[Context](interceptor: ServerInterceptor[Future, Future
1919
override def onMessage(message: Request): Unit = {
2020
import scala.concurrent.ExecutionContext.Implicits.global
2121
val responseMetadata = new Metadata()
22-
val futureResponse = interceptor.unary(ctx => logic(message, ctx))(RequestResponseMetadata(headers, responseMetadata))
22+
val futureResponse =
23+
interceptor.unary(message, ctx => logic(message, ctx))(using rpc.requestCodec, rpc.responseCodec)(
24+
RequestResponseMetadata(headers, responseMetadata)
25+
)
2326
futureResponse.onComplete { result =>
2427
result match {
2528
case scala.util.Success(response) =>
Lines changed: 34 additions & 7 deletions
Original file line numberDiff line numberDiff line change
@@ -1,14 +1,41 @@
11
package proteus.server
22

3-
trait ServerInterceptor[Unary[_], Streaming[_], InitialContext, Context] {
4-
def unary[A](io: Context => Unary[A]): (InitialContext => Unary[A])
5-
def stream[A](io: Context => Streaming[A]): (InitialContext => Streaming[A])
3+
import proteus.ProtobufCodec
4+
5+
trait ServerInterceptor[InitialUnary[_], Unary[_], InitialStreaming[_], Streaming[_], InitialContext, Context] {
6+
def unary[Req: ProtobufCodec, Resp: ProtobufCodec](request: Req, io: Context => Unary[Resp]): (InitialContext => InitialUnary[Resp])
7+
def clientStreaming[Req: ProtobufCodec, Resp: ProtobufCodec](
8+
io: Streaming[Req] => Context => Unary[Resp]
9+
): (InitialStreaming[Req] => InitialContext => InitialUnary[Resp])
10+
def serverStreaming[Req: ProtobufCodec, Resp: ProtobufCodec](
11+
request: Req,
12+
io: Context => Streaming[Resp]
13+
): (InitialContext => InitialStreaming[Resp])
14+
def bidiStreaming[Req: ProtobufCodec, Resp: ProtobufCodec](
15+
io: Streaming[Req] => Context => Streaming[Resp]
16+
): (InitialStreaming[Req] => InitialContext => InitialStreaming[Resp])
17+
}
18+
19+
trait ServerContextInterceptor[Unary[_], Streaming[_], InitialContext, Context]
20+
extends ServerInterceptor[Unary, Unary, Streaming, Streaming, InitialContext, Context] {
21+
def transformContext(context: InitialContext): Context
22+
def unary[Req: ProtobufCodec, Resp: ProtobufCodec](request: Req, io: Context => Unary[Resp]): (InitialContext => Unary[Resp]) =
23+
ctx => io(transformContext(ctx))
24+
def clientStreaming[Req: ProtobufCodec, Resp: ProtobufCodec](
25+
io: Streaming[Req] => Context => Unary[Resp]
26+
): (Streaming[Req] => InitialContext => Unary[Resp]) = stream => ctx => io(stream)(transformContext(ctx))
27+
def serverStreaming[Req: ProtobufCodec, Resp: ProtobufCodec](
28+
request: Req,
29+
io: Context => Streaming[Resp]
30+
): (InitialContext => Streaming[Resp]) = ctx => io(transformContext(ctx))
31+
def bidiStreaming[Req: ProtobufCodec, Resp: ProtobufCodec](
32+
io: Streaming[Req] => Context => Streaming[Resp]
33+
): (Streaming[Req] => InitialContext => Streaming[Resp]) = stream => ctx => io(stream)(transformContext(ctx))
634
}
735

836
object ServerInterceptor {
9-
def empty[Unary[_], Streaming[_], Context]: ServerInterceptor[Unary, Streaming, Context, Context] =
10-
new ServerInterceptor[Unary, Streaming, Context, Context] {
11-
def unary[A](io: Context => Unary[A]): (Context => Unary[A]) = io
12-
def stream[A](io: Context => Streaming[A]): (Context => Streaming[A]) = io
37+
def empty[Unary[_], Streaming[_], Context]: ServerContextInterceptor[Unary, Streaming, Context, Context] =
38+
new ServerContextInterceptor[Unary, Streaming, Context, Context] {
39+
def transformContext(context: Context): Context = context
1340
}
1441
}

0 commit comments

Comments
 (0)