2121using Grpc . Core ;
2222using Grpc . Core . Interceptors ;
2323
24- namespace Apache . Arrow . Flight . Middleware . Interceptors
24+ namespace Apache . Arrow . Flight . Middleware . Interceptors ;
25+
26+ public sealed class ClientInterceptorAdapter : Interceptor
2527{
26- public sealed class ClientInterceptorAdapter : Interceptor
28+ private readonly IReadOnlyList < IFlightClientMiddlewareFactory > _factories ;
29+
30+ public ClientInterceptorAdapter ( IEnumerable < IFlightClientMiddlewareFactory > factories )
2731 {
28- private readonly IReadOnlyList < IFlightClientMiddlewareFactory > _factories ;
32+ _factories = factories ? . ToList ( ) ?? throw new ArgumentNullException ( nameof ( factories ) ) ;
33+ }
2934
30- public ClientInterceptorAdapter ( IEnumerable < IFlightClientMiddlewareFactory > factories )
31- {
32- _factories = factories ? . ToList ( ) ?? throw new ArgumentNullException ( nameof ( factories ) ) ;
33- }
35+ public override AsyncUnaryCall < TResponse > AsyncUnaryCall < TRequest , TResponse > (
36+ TRequest request ,
37+ ClientInterceptorContext < TRequest , TResponse > context ,
38+ AsyncUnaryCallContinuation < TRequest , TResponse > continuation )
39+ where TRequest : class
40+ where TResponse : class
41+ {
42+ var options = InterceptCall ( context , out var middlewares ) ;
43+
44+ var newContext = new ClientInterceptorContext < TRequest , TResponse > (
45+ context . Method ,
46+ context . Host ,
47+ options ) ;
48+
49+ var call = continuation ( request , newContext ) ;
50+
51+ return new AsyncUnaryCall < TResponse > (
52+ HandleResponse ( call . ResponseAsync , call . ResponseHeadersAsync , call . GetStatus , call . GetTrailers ,
53+ call . Dispose , middlewares ) ,
54+ call . ResponseHeadersAsync ,
55+ call . GetStatus ,
56+ call . GetTrailers ,
57+ call . Dispose
58+ ) ;
59+ }
3460
35- public override AsyncUnaryCall < TResponse > AsyncUnaryCall < TRequest , TResponse > (
36- TRequest request ,
37- ClientInterceptorContext < TRequest , TResponse > context ,
38- AsyncUnaryCallContinuation < TRequest , TResponse > continuation )
39- where TRequest : class
40- where TResponse : class
41- {
42- var options = InterceptCall ( context , out var middlewares ) ;
43-
44- var newContext = new ClientInterceptorContext < TRequest , TResponse > (
45- context . Method ,
46- context . Host ,
47- options ) ;
48-
49- var call = continuation ( request , newContext ) ;
50-
51- return new AsyncUnaryCall < TResponse > (
52- HandleResponse ( call . ResponseAsync , call . ResponseHeadersAsync , call . GetStatus , call . GetTrailers ,
53- call . Dispose , middlewares ) ,
54- call . ResponseHeadersAsync ,
55- call . GetStatus ,
56- call . GetTrailers ,
57- call . Dispose
58- ) ;
59- }
61+ public override AsyncServerStreamingCall < TResponse > AsyncServerStreamingCall < TRequest , TResponse > (
62+ TRequest request ,
63+ ClientInterceptorContext < TRequest , TResponse > context ,
64+ AsyncServerStreamingCallContinuation < TRequest , TResponse > continuation )
65+ where TRequest : class
66+ where TResponse : class
67+ {
68+ var callOptions = InterceptCall ( context , out var middlewares ) ;
69+ var newContext = new ClientInterceptorContext < TRequest , TResponse > (
70+ context . Method , context . Host , callOptions ) ;
6071
61- public override AsyncServerStreamingCall < TResponse > AsyncServerStreamingCall < TRequest , TResponse > (
62- TRequest request ,
63- ClientInterceptorContext < TRequest , TResponse > context ,
64- AsyncServerStreamingCallContinuation < TRequest , TResponse > continuation )
65- where TRequest : class
66- where TResponse : class
72+ var call = continuation ( request , newContext ) ;
73+
74+ var responseHeadersTask = call . ResponseHeadersAsync . ContinueWith ( task =>
6775 {
68- var callOptions = InterceptCall ( context , out var middlewares ) ;
69- var newContext = new ClientInterceptorContext < TRequest , TResponse > (
70- context . Method , context . Host , callOptions ) ;
76+ if ( task . IsFaulted )
77+ {
78+ throw task . Exception ! ;
79+ }
80+
81+ if ( task . IsCanceled )
82+ {
83+ throw new TaskCanceledException ( task ) ;
84+ }
85+
86+ var headers = task . Result ;
87+ var ch = new CallHeaders ( headers ) ;
88+ foreach ( var m in middlewares )
89+ m ? . OnHeadersReceived ( ch ) ;
90+
91+ return headers ;
92+ } ) ;
93+
94+ var wrappedResponseStream = new MiddlewareResponseStream < TResponse > (
95+ call . ResponseStream ,
96+ call ,
97+ middlewares ) ;
98+
99+ return new AsyncServerStreamingCall < TResponse > (
100+ wrappedResponseStream ,
101+ responseHeadersTask ,
102+ call . GetStatus ,
103+ call . GetTrailers ,
104+ call . Dispose ) ;
105+ }
71106
72- var call = continuation ( request , newContext ) ;
73107
74- var responseHeadersTask = call . ResponseHeadersAsync . ContinueWith ( task =>
75- {
76- if ( task . Exception == null && task . Result != null )
77- {
78- var headers = task . Result ;
79- foreach ( var m in middlewares )
80- m ? . OnHeadersReceived ( new CallHeaders ( headers ) ) ;
81- }
82-
83- return task . Result ;
84- } ) ;
85-
86- var wrappedResponseStream = new MiddlewareResponseStream < TResponse > (
87- call . ResponseStream ,
88- call ,
89- middlewares ) ;
90-
91- return new AsyncServerStreamingCall < TResponse > (
92- wrappedResponseStream ,
93- responseHeadersTask ,
94- call . GetStatus ,
95- call . GetTrailers ,
96- call . Dispose ) ;
97- }
108+ private CallOptions InterceptCall < TRequest , TResponse > (
109+ ClientInterceptorContext < TRequest , TResponse > context ,
110+ out List < IFlightClientMiddleware > middlewareList )
111+ where TRequest : class
112+ where TResponse : class
113+ {
114+ var callInfo = new CallInfo ( context . Method . FullName , context . Method . Type ) ;
98115
116+ var headers = context . Options . Headers ?? new Metadata ( ) ;
117+ middlewareList = new List < IFlightClientMiddleware > ( ) ;
99118
100- private CallOptions InterceptCall < TRequest , TResponse > (
101- ClientInterceptorContext < TRequest , TResponse > context ,
102- out List < IFlightClientMiddleware > middlewareList )
103- where TRequest : class
104- where TResponse : class
105- {
106- var callInfo = new CallInfo ( context . Method . FullName , context . Method . Type ) ;
119+ var callHeaders = new CallHeaders ( headers ) ;
107120
108- var headers = context . Options . Headers ?? new Metadata ( ) ;
109- middlewareList = new List < IFlightClientMiddleware > ( ) ;
121+ foreach ( var factory in _factories )
122+ {
123+ var middleware = factory . OnCallStarted ( callInfo ) ;
124+ middleware ? . OnBeforeSendingHeaders ( callHeaders ) ;
125+ middlewareList . Add ( middleware ) ;
126+ }
110127
111- var callHeaders = new CallHeaders ( headers ) ;
128+ return context . Options . WithHeaders ( headers ) ;
129+ }
112130
113- foreach ( var factory in _factories )
131+ private async Task < TResponse > HandleResponse < TResponse > (
132+ Task < TResponse > responseTask ,
133+ Task < Metadata > headersTask ,
134+ Func < Status > getStatus ,
135+ Func < Metadata > getTrailers ,
136+ Action dispose ,
137+ List < IFlightClientMiddleware > middlewares )
138+ {
139+ var nonNullMiddlewares = ( middlewares ?? new List < IFlightClientMiddleware > ( ) )
140+ . Where ( m => m != null )
141+ . ToList ( ) ;
142+
143+ var hasMiddlewares = nonNullMiddlewares . Count > 0 ;
144+ var completionNotified = false ;
145+
146+ try
147+ {
148+ // Always await headers to surface faults; only materialize CallHeaders if needed.
149+ var headers = await headersTask . ConfigureAwait ( false ) ;
150+ if ( hasMiddlewares )
114151 {
115- var middleware = factory . OnCallStarted ( callInfo ) ;
116- middleware ? . OnBeforeSendingHeaders ( callHeaders ) ;
117- middlewareList . Add ( middleware ) ;
152+ var ch = new CallHeaders ( headers ) ;
153+ foreach ( var m in nonNullMiddlewares )
154+ m . OnHeadersReceived ( ch ) ;
118155 }
119156
120- return context . Options . WithHeaders ( headers ) ;
121- }
157+ var response = await responseTask . ConfigureAwait ( false ) ;
122158
123- private async Task < TResponse > HandleResponse < TResponse > (
124- Task < TResponse > responseTask ,
125- Task < Metadata > headersTask ,
126- Func < Status > getStatus ,
127- Func < Metadata > getTrailers ,
128- Action dispose ,
129- List < IFlightClientMiddleware > middlewares )
159+ // Single completion notification
160+ NotifyCompletionOnce ( ) ;
161+ return response ;
162+ }
163+ catch
130164 {
131- try
132- {
133- var headers = await headersTask . ConfigureAwait ( false ) ;
134- foreach ( var m in middlewares )
135- {
136- m ? . OnHeadersReceived ( new CallHeaders ( headers ) ) ;
137- }
138-
139- var response = await responseTask . ConfigureAwait ( false ) ;
140- foreach ( var m in middlewares )
141- {
142- m ? . OnCallCompleted ( getStatus ( ) , getTrailers ( ) ) ;
143- }
144-
145- return response ;
146- }
147- catch
148- {
149- foreach ( var m in middlewares )
150- {
151- m ? . OnCallCompleted ( getStatus ( ) , getTrailers ( ) ) ;
152- }
153- throw ;
154- }
155- finally
156- {
157- dispose ? . Invoke ( ) ;
158- }
165+ // Completion on failure (only once)
166+ NotifyCompletionOnce ( ) ;
167+ throw ;
168+ }
169+ finally
170+ {
171+ dispose ? . Invoke ( ) ;
172+ }
173+
174+ void NotifyCompletionOnce ( )
175+ {
176+ if ( completionNotified || ! hasMiddlewares ) return ;
177+ completionNotified = true ;
178+
179+ var status = getStatus ( ) ;
180+ var trailers = getTrailers ( ) ;
181+
182+ foreach ( var m in nonNullMiddlewares )
183+ m . OnCallCompleted ( status , trailers ) ;
159184 }
160185 }
161186}
0 commit comments