Skip to content

Commit 1605fad

Browse files
committed
Refactor ClientInterceptorAdapter to reduce duplication and improve middleware handling
- Switched to file-scoped namespace style for consistency - Ensured exceptions from ResponseHeadersAsync are explicitly propagated - Lifted CallHeaders allocation out of loop to avoid redundant allocations - Captured getStatus() and getTrailers() once per call instead of per-middleware - Fixed double invocation of OnCallCompleted when a middleware throws - Extracted NotifyCompletionOnce as a private helper to keep HandleResponse DRY - Preserved behavior while ensuring OnCallCompleted is invoked at most once
1 parent 67360a1 commit 1605fad

5 files changed

Lines changed: 155 additions & 238 deletions

File tree

csharp/src/Apache.Arrow.Flight/Middleware/ClientCookieMiddleware.cs

Lines changed: 6 additions & 9 deletions
Original file line numberDiff line numberDiff line change
@@ -24,8 +24,8 @@ public class ClientCookieMiddleware : IFlightClientMiddleware
2424
{
2525
private readonly ClientCookieMiddlewareFactory _factory;
2626
private readonly ILogger<ClientCookieMiddleware> _logger;
27-
private const string SET_COOKIE_HEADER = "Set-Cookie";
28-
private const string COOKIE_HEADER = "Cookie";
27+
private const string SetCookieHeader = "Set-Cookie";
28+
private const string CookieHeader = "Cookie";
2929

3030
public ClientCookieMiddleware(ClientCookieMiddlewareFactory factory,
3131
ILogger<ClientCookieMiddleware> logger)
@@ -41,29 +41,26 @@ public void OnBeforeSendingHeaders(ICallHeaders outgoingHeaders)
4141
var cookieValue = GetValidCookiesAsString();
4242
if (!string.IsNullOrEmpty(cookieValue))
4343
{
44-
outgoingHeaders.Insert(COOKIE_HEADER, cookieValue);
44+
outgoingHeaders.Insert(CookieHeader, cookieValue);
4545
}
46-
_logger.LogInformation("Sending Headers: " + string.Join(", ", outgoingHeaders));
4746
}
4847

4948
public void OnHeadersReceived(ICallHeaders incomingHeaders)
5049
{
51-
var setCookies = incomingHeaders.GetAll(SET_COOKIE_HEADER);
50+
var setCookies = incomingHeaders.GetAll(SetCookieHeader);
5251
_factory.UpdateCookies(setCookies);
53-
_logger.LogInformation("Received Headers: " + string.Join(", ", incomingHeaders));
5452
}
5553

5654
public void OnCallCompleted(Status status, Metadata trailers)
5755
{
58-
_logger.LogInformation($"Call completed with: {status.StatusCode} ({status.Detail})");
56+
// ingest: status and/or metadata trailers
5957
}
6058

6159
private string GetValidCookiesAsString()
6260
{
6361
var cookieList = new List<string>();
6462
foreach (var entry in _factory.Cookies)
6563
{
66-
_logger.LogInformation($"Before remove cookie: {entry.Key} Expired: ({entry.Value.Expired})");
6764
if (entry.Value.Expired)
6865
{
6966
_factory.Cookies.TryRemove(entry.Key, out _);
@@ -75,4 +72,4 @@ private string GetValidCookiesAsString()
7572
}
7673
return string.Join("; ", cookieList);
7774
}
78-
}
75+
}

csharp/src/Apache.Arrow.Flight/Middleware/ClientCookieMiddlewareFactory.cs

Lines changed: 6 additions & 9 deletions
Original file line numberDiff line numberDiff line change
@@ -21,27 +21,26 @@
2121
using Apache.Arrow.Flight.Middleware.Extensions;
2222
using Apache.Arrow.Flight.Middleware.Interfaces;
2323
using Microsoft.Extensions.Logging;
24+
2425
namespace Apache.Arrow.Flight.Middleware;
2526

2627
public class ClientCookieMiddlewareFactory : IFlightClientMiddlewareFactory
2728
{
2829
public readonly ConcurrentDictionary<string, Cookie> Cookies = new(StringComparer.OrdinalIgnoreCase);
29-
private readonly ILoggerFactory _loggerFactory;
30+
private readonly ILogger<ClientCookieMiddleware> _logger;
3031

3132
public ClientCookieMiddlewareFactory(ILoggerFactory loggerFactory)
3233
{
33-
_loggerFactory = loggerFactory;
34+
_logger = loggerFactory.CreateLogger<ClientCookieMiddleware>();
3435
}
3536

3637
public IFlightClientMiddleware OnCallStarted(CallInfo callInfo)
3738
{
38-
var logger = _loggerFactory.CreateLogger<ClientCookieMiddleware>();
39-
return new ClientCookieMiddleware(this, logger);
39+
return new ClientCookieMiddleware(this, _logger);
4040
}
41-
41+
4242
internal void UpdateCookies(IEnumerable<string> newCookieHeaderValues)
4343
{
44-
var logger = _loggerFactory.CreateLogger<ClientCookieMiddleware>();
4544
foreach (var headerValue in newCookieHeaderValues)
4645
{
4746
try
@@ -61,10 +60,8 @@ internal void UpdateCookies(IEnumerable<string> newCookieHeaderValues)
6160
}
6261
catch (FormatException ex)
6362
{
64-
65-
logger.LogWarning(ex, "Skipping malformed Set-Cookie header: '{HeaderValue}'", headerValue);
63+
_logger.LogWarning(ex, "Skipping malformed Set-Cookie header: '{HeaderValue}'", headerValue);
6664
}
6765
}
6866
}
69-
7067
}

csharp/src/Apache.Arrow.Flight/Middleware/Interceptors/ClientInterceptorAdapter.cs

Lines changed: 142 additions & 117 deletions
Original file line numberDiff line numberDiff line change
@@ -21,141 +21,166 @@
2121
using Grpc.Core;
2222
using 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
}

csharp/src/Apache.Arrow.Flight/Middleware/Interfaces/IFlightClientMiddlewareFactory.cs

Lines changed: 1 addition & 1 deletion
Original file line numberDiff line numberDiff line change
@@ -18,4 +18,4 @@ namespace Apache.Arrow.Flight.Middleware.Interfaces;
1818
public interface IFlightClientMiddlewareFactory
1919
{
2020
IFlightClientMiddleware OnCallStarted(CallInfo callInfo);
21-
}
21+
}

0 commit comments

Comments
 (0)