Skip to content

Commit 7f2ce1d

Browse files
committed
feat(sdk): ProviderDrain 优雅下线对齐(go/js/java/cpp)+ Java TLS 传输
- 收到 ProviderDrainRequest 置位 draining(拒绝新 Invoke)、在途调用 清零后复用既有重连编排恢复会话,立即回空确认帧(幂等) - go: tcpRPCHandler.handleDrain + atomic 状态;js: BasicClient 对齐 C# 参考实现;java: 协议消息 + client 接线 + TlsSocketFactory - 新增测试:go drain_test、js drain.test、java Drain/TlsTransport 测试 验证:go test(drain 用例)与 js jest(drain 用例)本机通过; java/cpp 本机无 mvn/cmake 未编译验证
1 parent fbc40d1 commit 7f2ce1d

15 files changed

Lines changed: 1031 additions & 2 deletions

File tree

sdks/cpp/include/croupier/sdk/croupier_client.h

Lines changed: 4 additions & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -246,6 +246,10 @@ class CroupierClient {
246246
// Check if the client is connected to the agent
247247
bool IsConnected() const;
248248

249+
// 是否处于 drain 状态(收到 Agent 的 ProviderDrainRequest 后为 true,
250+
// 恢复完成或停止后为 false)
251+
bool IsDraining() const;
252+
249253
// Start serving (blocking call until Stop() is called)
250254
void Serve();
251255

sdks/cpp/include/croupier/sdk/tcp_transport.h

Lines changed: 1 addition & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -248,6 +248,7 @@ class TCPServer {
248248
*/
249249
std::string GetListenAddress() const;
250250

251+
251252
private:
252253
struct ClientConnection {
253254
socket_t socket;

sdks/cpp/src/croupier_client.cpp

Lines changed: 55 additions & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -338,6 +338,11 @@ class CroupierClient::Impl {
338338
std::optional<std::thread::id> heartbeat_thread_id_;
339339
std::string last_error_;
340340

341+
// Drain 状态:收到 ProviderDrainRequest 后置位——拒绝新 Invoke,
342+
// 在途调用清零后复用 RegisterAllFunctions 恢复会话(对齐 C# 参考实现)。
343+
std::atomic<bool> draining_{false};
344+
std::atomic<int64_t> inflight_calls_{0};
345+
341346
// Reconnection state
342347
std::atomic<bool> is_reconnecting_{false};
343348
std::atomic<bool> should_stop_reconnecting_{false};
@@ -423,11 +428,57 @@ class CroupierClient::Impl {
423428
}
424429
}
425430

431+
// 在途调用 RAII 计数:drain 恢复以其清零为信号。
432+
struct InflightGuard {
433+
explicit InflightGuard(Impl* impl) : impl_(impl) { impl_->inflight_calls_.fetch_add(1); }
434+
~InflightGuard() { impl_->inflight_calls_.fetch_sub(1); }
435+
Impl* impl_;
436+
};
437+
438+
// 等待在途调用完成(最多 30s),随后按 auto_reconnect 语义恢复会话。
439+
void DrainAndRecover() {
440+
const auto deadline = std::chrono::steady_clock::now() + std::chrono::seconds(30);
441+
while (inflight_calls_.load() > 0 && std::chrono::steady_clock::now() < deadline
442+
&& running_.load()) {
443+
std::this_thread::sleep_for(std::chrono::milliseconds(100));
444+
}
445+
if (inflight_calls_.load() > 0) {
446+
SDK_LOG_ERROR("Drain timeout with in-flight calls still running");
447+
}
448+
if (config_.auto_reconnect && running_.load()) {
449+
SDK_LOG_INFO("Drain complete, reconnecting provider session");
450+
try {
451+
RegisterAllFunctions();
452+
} catch (const std::exception& e) {
453+
SDK_LOG_ERROR(std::string("Drain recovery failed: ") + e.what());
454+
} catch (...) {
455+
SDK_LOG_ERROR("Drain recovery failed: unknown error");
456+
}
457+
}
458+
draining_.store(false);
459+
}
460+
426461
// handleAgentRequest 处理 Agent -> Provider 调用(invoke / start task),
427462
// 由 TCPTransport 的有界 worker 池并发执行(读循环只投递)。
428463
std::vector<uint8_t> handleAgentRequest(uint32_t msg_id, uint32_t /*req_id*/, const std::vector<uint8_t>& body) {
429464
try {
465+
if (msg_id == protocol::MSG_PROVIDER_DRAIN_REQUEST) {
466+
// 幂等:重复 drain 只回确认。置位后异步等待在途清零再恢复。
467+
if (!draining_.exchange(true)) {
468+
SDK_LOG_INFO("Drain requested");
469+
std::thread([this]() { DrainAndRecover(); }).detach();
470+
}
471+
::croupier::sdk::v1::ProviderDrainResponse resp;
472+
return SerializeMessage(resp);
473+
}
430474
if (msg_id == protocol::MSG_INVOKE_REQUEST) {
475+
// drain 期间拒绝新调用,等待 Agent 停止投递。
476+
if (draining_.load()) {
477+
::croupier::sdk::v1::InvokeResponse resp;
478+
resp.set_payload("{\"error\":\"provider is draining\"}");
479+
return SerializeMessage(resp);
480+
}
481+
InflightGuard guard(this);
431482
auto req = ParseMessage<::croupier::sdk::v1::InvokeRequest>(body, "InvokeRequest");
432483
auto it = handlers_.find(req.function_id());
433484
if (it == handlers_.end()) {
@@ -2074,6 +2125,10 @@ bool CroupierClient::IsConnected() const {
20742125
return impl_->IsConnected();
20752126
}
20762127

2128+
bool CroupierClient::IsDraining() const {
2129+
return impl_->draining_.load();
2130+
}
2131+
20772132
void CroupierClient::Serve() {
20782133
impl_->Serve();
20792134
}

sdks/cpp/tests/test_provider_inbound.cpp

Lines changed: 97 additions & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -105,6 +105,13 @@ class RawFakeAgent {
105105
WriteFrame(msg_id, req_id, body);
106106
}
107107

108+
// 接受客户端的重连(drain 恢复 / 断线重连均会发起二次握手)。
109+
void AcceptAndHandshakeAgain() {
110+
if (conn_ != INVALID_SOCK) closesocket(conn_);
111+
conn_ = INVALID_SOCK;
112+
AcceptAndHandshake();
113+
}
114+
108115
protocol::ParsedMessage ReadFrame() {
109116
uint8_t hdr[4] = {0};
110117
if (!ReadAll(conn_, hdr, 4)) ADD_FAILURE() << "read frame header failed";
@@ -332,3 +339,93 @@ TEST(ClientAddressTest, HTTPSchemeRejectedForTCPClient) {
332339

333340
} // namespace
334341
} // namespace croupier::sdk::test
342+
343+
namespace croupier::sdk::test {
344+
namespace {
345+
346+
// Agent 下发 drain:客户端立即回空确认并置位 IsDraining;
347+
// drain 期间新 Invoke 被拒(错误 payload);恢复后状态清除。
348+
TEST(ProviderInboundTest, AgentDrainAcksRejectsInvokeAndRecovers) {
349+
RawFakeAgent agent;
350+
std::thread agent_thread([&] { agent.AcceptAndHandshake(); });
351+
352+
CroupierClient client(ProviderConfig(agent.address()));
353+
FunctionDescriptor desc;
354+
desc.id = "test.echo";
355+
desc.version = "1.0.0";
356+
desc.operation = "echo";
357+
desc.capability = "action";
358+
desc.risk = "safe";
359+
std::atomic<int> calls{0};
360+
ASSERT_TRUE(client.RegisterFunction(
361+
desc, [&](const std::string&, const std::string& payload) {
362+
calls.fetch_add(1);
363+
std::this_thread::sleep_for(std::chrono::milliseconds(120));
364+
return "echo:" + payload;
365+
}));
366+
ASSERT_TRUE(client.Connect());
367+
agent_thread.join();
368+
369+
EXPECT_FALSE(client.IsDraining());
370+
371+
// 在途调用先行:handler 睡 120ms,drain 必须等它完成
372+
agent.PushRequest(protocol::MSG_INVOKE_REQUEST, 9101, InvokeBody("test.echo", "inflight"));
373+
std::this_thread::sleep_for(std::chrono::milliseconds(30));
374+
375+
// 推 drain 请求(req_id 9102)
376+
agent.PushRequest(protocol::MSG_PROVIDER_DRAIN_REQUEST, 9102, {});
377+
378+
// drain 期间的新 Invoke 被拒(handler 不应执行)
379+
agent.PushRequest(protocol::MSG_INVOKE_REQUEST, 9103, InvokeBody("test.echo", "rejected"));
380+
auto rejected = agent.ReadResponseFor(9103);
381+
v1::InvokeResponse rejected_resp;
382+
ASSERT_TRUE(rejected_resp.ParseFromArray(rejected.body.data(), static_cast<int>(rejected.body.size())));
383+
EXPECT_NE(rejected_resp.payload().find("provider is draining"), std::string::npos);
384+
EXPECT_EQ(calls.load(), 1); // 只有在途那次执行了
385+
386+
// drain 确认帧(空 ProviderDrainResponse)
387+
auto ack = agent.ReadResponseFor(9102);
388+
EXPECT_EQ(ack.msg_id, protocol::MSG_PROVIDER_DRAIN_RESPONSE);
389+
EXPECT_TRUE(ack.body.empty());
390+
391+
// 等在途完成 + 恢复(auto_reconnect 默认 true → 重连重注册)
392+
std::thread reconnect_thread([&] { agent.AcceptAndHandshakeAgain(); });
393+
reconnect_thread.join();
394+
for (int i = 0; i < 50 && client.IsDraining(); ++i) {
395+
std::this_thread::sleep_for(std::chrono::milliseconds(20));
396+
}
397+
EXPECT_FALSE(client.IsDraining());
398+
399+
client.Close();
400+
}
401+
402+
// drain 幂等:重复请求只回确认,不重复触发恢复。
403+
TEST(ProviderInboundTest, AgentDrainIsIdempotent) {
404+
RawFakeAgent agent;
405+
std::thread agent_thread([&] { agent.AcceptAndHandshake(); });
406+
407+
CroupierClient client(ProviderConfig(agent.address()));
408+
FunctionDescriptor desc;
409+
desc.id = "test.echo";
410+
desc.version = "1.0.0";
411+
desc.operation = "echo";
412+
desc.capability = "action";
413+
desc.risk = "safe";
414+
ASSERT_TRUE(client.RegisterFunction(
415+
desc, [](const std::string&, const std::string& payload) { return payload; }));
416+
ASSERT_TRUE(client.Connect());
417+
agent_thread.join();
418+
419+
agent.PushRequest(protocol::MSG_PROVIDER_DRAIN_REQUEST, 9201, {});
420+
auto ack1 = agent.ReadResponseFor(9201);
421+
EXPECT_EQ(ack1.msg_id, protocol::MSG_PROVIDER_DRAIN_RESPONSE);
422+
agent.PushRequest(protocol::MSG_PROVIDER_DRAIN_REQUEST, 9202, {});
423+
auto ack2 = agent.ReadResponseFor(9202);
424+
EXPECT_EQ(ack2.msg_id, protocol::MSG_PROVIDER_DRAIN_RESPONSE);
425+
EXPECT_TRUE(client.IsDraining());
426+
427+
client.Close();
428+
}
429+
430+
} // namespace
431+
} // namespace croupier::sdk::test

sdks/go/pkg/croupier/drain_test.go

Lines changed: 135 additions & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -0,0 +1,135 @@
1+
// Copyright 2025 Croupier Authors
2+
// Licensed under the Apache License, Version 2.0
3+
4+
package croupier
5+
6+
import (
7+
"context"
8+
"sync/atomic"
9+
"testing"
10+
"time"
11+
12+
sdkv1 "github.com/cuihairu/croupier/sdks/go/pkg/pb/croupier/sdk/v1"
13+
"google.golang.org/protobuf/proto"
14+
15+
"github.com/cuihairu/croupier/sdks/go/pkg/croupier/protocol"
16+
)
17+
18+
// drain 请求:置位状态、立即回空确认、幂等;drain 期间新 Invoke 被拒。
19+
func TestDrainHandler_AcksIdempotentAndRejectsInvoke(t *testing.T) {
20+
handler := newTestRPCHandler(t)
21+
handler.manager.handlers["test.fn"] = func(ctx context.Context, payload []byte) ([]byte, error) {
22+
return []byte(`{}`), nil
23+
}
24+
25+
drainReq, _ := proto.Marshal(&sdkv1.ProviderDrainRequest{SessionId: "s-1", Reason: "rolling-restart", RetryAfterMs: 1000})
26+
respBody, err := handler.handleDrain(context.Background(), protocol.MsgProviderDrainRequest, 1, drainReq)
27+
if err != nil {
28+
t.Fatalf("handleDrain: %v", err)
29+
}
30+
if err := proto.Unmarshal(respBody, &sdkv1.ProviderDrainResponse{}); err != nil {
31+
t.Fatalf("response must be ProviderDrainResponse: %v", err)
32+
}
33+
if !handler.manager.draining.Load() {
34+
t.Fatal("draining must be set after first drain request")
35+
}
36+
37+
// 幂等:重复 drain 不重复触发恢复
38+
if _, err := handler.handleDrain(context.Background(), protocol.MsgProviderDrainRequest, 2, drainReq); err != nil {
39+
t.Fatalf("idempotent drain: %v", err)
40+
}
41+
42+
// drain 期间新 Invoke 被拒:返回 provider is draining 错误 payload,handler 不执行
43+
called := false
44+
handler.manager.handlers["test.fn"] = func(ctx context.Context, payload []byte) ([]byte, error) {
45+
called = true
46+
return []byte(`{}`), nil
47+
}
48+
invokeReq, _ := proto.Marshal(&sdkv1.InvokeRequest{FunctionId: "test.fn"})
49+
respBody, err = handler.invoke(context.Background(), protocol.MsgInvokeRequest, 3, invokeReq)
50+
if err != nil {
51+
t.Fatalf("invoke during drain must not be transport error: %v", err)
52+
}
53+
if called {
54+
t.Fatal("handler must not run while draining")
55+
}
56+
resp := &sdkv1.InvokeResponse{}
57+
if err := proto.Unmarshal(respBody, resp); err != nil {
58+
t.Fatalf("unmarshal: %v", err)
59+
}
60+
if string(resp.Payload) != `{"error":"provider is draining"}` {
61+
t.Fatalf("unexpected draining payload: %s", resp.Payload)
62+
}
63+
64+
// 恢复等待在途为 0 后 handleDisconnect:无 onDisconnect、Reconnect nil → 断开
65+
if handler.manager.config.Reconnect != nil && handler.manager.config.Reconnect.Enabled {
66+
handler.manager.config.Reconnect.Enabled = false
67+
}
68+
handler.manager.drainAndRecover()
69+
if handler.manager.draining.Load() {
70+
t.Fatal("draining must clear after recovery")
71+
}
72+
}
73+
74+
// 在途计数:invoke 期间 inflight>0,drainAndRecover 等待其清零。
75+
func TestDrain_WaitsForInflightCalls(t *testing.T) {
76+
handler := newTestRPCHandler(t)
77+
handler.manager.config.Reconnect = nil
78+
handler.manager.handlers["slow.fn"] = func(ctx context.Context, payload []byte) ([]byte, error) {
79+
time.Sleep(200 * time.Millisecond)
80+
return []byte(`{}`), nil
81+
}
82+
83+
done := make(chan struct{})
84+
go func() {
85+
defer close(done)
86+
_, _ = handler.invoke(context.Background(), protocol.MsgInvokeRequest, 1, mustMarshal(t, &sdkv1.InvokeRequest{FunctionId: "slow.fn"}))
87+
}()
88+
time.Sleep(50 * time.Millisecond) // 等 invoke 进入 handler
89+
if n := handler.manager.inflightCalls.Load(); n != 1 {
90+
t.Fatalf("inflight = %d, want 1", n)
91+
}
92+
93+
recovered := make(chan struct{})
94+
go func() {
95+
handler.manager.draining.Store(true)
96+
handler.manager.drainAndRecover()
97+
close(recovered)
98+
}()
99+
select {
100+
case <-recovered:
101+
t.Fatal("drainAndRecover must wait for in-flight call")
102+
case <-time.After(100 * time.Millisecond):
103+
}
104+
<-done
105+
select {
106+
case <-recovered:
107+
case <-time.After(time.Second):
108+
t.Fatal("drainAndRecover did not finish after in-flight completed")
109+
}
110+
}
111+
112+
// drain 后经 handleDisconnect 触发 onDisconnect(对齐既有重连编排入口)。
113+
func TestDrain_FiresOnDisconnectForReconnect(t *testing.T) {
114+
handler := newTestRPCHandler(t)
115+
handler.manager.config.Reconnect = &ReconnectConfig{Enabled: true}
116+
var fired atomic.Bool
117+
handler.manager.onDisconnect = func() { fired.Store(true) }
118+
handler.manager.connected = true
119+
handler.manager.draining.Store(true)
120+
121+
handler.manager.drainAndRecover()
122+
123+
if !fired.Load() {
124+
t.Fatal("onDisconnect must fire so client.go reconnect loop takes over")
125+
}
126+
}
127+
128+
func mustMarshal(t *testing.T, m proto.Message) []byte {
129+
t.Helper()
130+
b, err := proto.Marshal(m)
131+
if err != nil {
132+
t.Fatalf("marshal: %v", err)
133+
}
134+
return b
135+
}

0 commit comments

Comments
 (0)