diff --git a/cmd/optional-driver-agent/main.go b/cmd/optional-driver-agent/main.go index 03d45e61..2af57caf 100644 --- a/cmd/optional-driver-agent/main.go +++ b/cmd/optional-driver-agent/main.go @@ -42,25 +42,28 @@ type agentResponse struct { } const ( - agentMethodConnect = "connect" - agentMethodClose = "close" - agentMethodMetadata = "metadata" - agentMethodPing = "ping" - agentMethodOpenSession = "openSession" - agentMethodCloseSession = "closeSession" - agentMethodQuery = "query" - agentMethodQueryMulti = "queryMulti" - agentMethodStreamQuery = "streamQuery" - agentMethodExec = "exec" - agentMethodGetDatabases = "getDatabases" - agentMethodGetTables = "getTables" - agentMethodGetCreateStmt = "getCreateStatement" - agentMethodGetColumns = "getColumns" - agentMethodGetAllColumns = "getAllColumns" - agentMethodGetIndexes = "getIndexes" - agentMethodGetForeignKey = "getForeignKeys" - agentMethodGetTriggers = "getTriggers" - agentMethodApplyChanges = "applyChanges" + agentMethodConnect = "connect" + agentMethodClose = "close" + agentMethodMetadata = "metadata" + agentMethodPing = "ping" + agentMethodOpenSession = "openSession" + agentMethodCloseSession = "closeSession" + agentMethodOpenTransaction = "openTransaction" + agentMethodCommitTransaction = "commitTransaction" + agentMethodRollbackTransaction = "rollbackTransaction" + agentMethodQuery = "query" + agentMethodQueryMulti = "queryMulti" + agentMethodStreamQuery = "streamQuery" + agentMethodExec = "exec" + agentMethodGetDatabases = "getDatabases" + agentMethodGetTables = "getTables" + agentMethodGetCreateStmt = "getCreateStatement" + agentMethodGetColumns = "getColumns" + agentMethodGetAllColumns = "getAllColumns" + agentMethodGetIndexes = "getIndexes" + agentMethodGetForeignKey = "getForeignKeys" + agentMethodGetTriggers = "getTriggers" + agentMethodApplyChanges = "applyChanges" ) const legacyClickHouseDefaultTimeout = 2 * time.Hour @@ -222,6 +225,23 @@ func handleRequest(runtimeState *agentRuntime, req agentRequest) agentResponse { runtimeState.sessions[sessionID] = session resp.Data = sessionID return resp + case agentMethodOpenTransaction: + if runtimeState.inst == nil { + return fail(resp, "connection not open") + } + provider, ok := runtimeState.inst.(db.TransactionExecerProvider) + if !ok { + return fail(resp, fmt.Sprintf("当前数据源(%s)不支持 SQL 编辑器托管事务", strings.TrimSpace(agentDriverType))) + } + // The transaction must outlive this request and be finished by a later RPC. + transaction, err := provider.OpenTransactionExecer(context.Background()) + if err != nil { + return fail(resp, err.Error()) + } + sessionID := runtimeState.nextID() + runtimeState.sessions[sessionID] = transaction + resp.Data = sessionID + return resp case agentMethodCloseSession: if err := runtimeState.closeSession(req.SessionID); err != nil { return fail(resp, err.Error()) @@ -261,6 +281,20 @@ func handleRequest(runtimeState *agentRuntime, req agentRequest) agentResponse { return fail(resp, err.Error()) } resp.RowsAffected = affected + case agentMethodCommitTransaction, agentMethodRollbackTransaction: + transaction, ok := session.(db.TransactionExecer) + if !ok { + return fail(resp, "当前会话不是托管事务") + } + var err error + if method == agentMethodCommitTransaction { + err = transaction.Commit() + } else { + err = transaction.Rollback() + } + if err != nil { + return fail(resp, err.Error()) + } default: return fail(resp, "当前事务会话不支持该方法") } diff --git a/cmd/optional-driver-agent/main_test.go b/cmd/optional-driver-agent/main_test.go index 5db34e92..70fbe953 100644 --- a/cmd/optional-driver-agent/main_test.go +++ b/cmd/optional-driver-agent/main_test.go @@ -176,6 +176,32 @@ type fakeAgentSessionDB struct { session *fakeAgentStatementSession } +type fakeAgentTransactionDB struct { + fakeAgentTimeoutDB + transaction *fakeAgentTransactionSession +} + +func (f *fakeAgentTransactionDB) OpenTransactionExecer(context.Context) (db.TransactionExecer, error) { + f.transaction = &fakeAgentTransactionSession{} + return f.transaction, nil +} + +type fakeAgentTransactionSession struct { + fakeAgentStatementSession + commitCalls int + rollbackCalls int +} + +func (f *fakeAgentTransactionSession) Commit() error { + f.commitCalls++ + return nil +} + +func (f *fakeAgentTransactionSession) Rollback() error { + f.rollbackCalls++ + return nil +} + func (f *fakeAgentSessionDB) OpenSessionExecer(ctx context.Context) (db.StatementExecer, error) { f.session = &fakeAgentStatementSession{} return f.session, nil @@ -468,6 +494,72 @@ func TestHandleRequest_UsesPinnedSessionForSessionScopedQueryAndExec(t *testing. } } +func TestHandleRequest_UsesManagedTransactionSession(t *testing.T) { + for _, tc := range []struct { + name string + finishMethod string + wantCommits int + wantRollbacks int + }{ + {name: "commit", finishMethod: "commitTransaction", wantCommits: 1}, + {name: "rollback", finishMethod: "rollbackTransaction", wantRollbacks: 1}, + } { + t.Run(tc.name, func(t *testing.T) { + fake := &fakeAgentTransactionDB{} + runtimeState := &agentRuntime{ + inst: fake, + sessions: make(map[string]db.StatementExecer), + } + + openResp := handleRequest(runtimeState, agentRequest{ID: 1, Method: "openTransaction"}) + if !openResp.Success { + t.Fatalf("openTransaction failed: %s", openResp.Error) + } + sessionID, ok := openResp.Data.(string) + if !ok || strings.TrimSpace(sessionID) == "" { + t.Fatalf("unexpected transaction id payload: %#v", openResp.Data) + } + + execResp := handleRequest(runtimeState, agentRequest{ + ID: 2, + Method: agentMethodExec, + SessionID: sessionID, + Query: "UPDATE t SET v = 1", + }) + if !execResp.Success { + t.Fatalf("transaction exec failed: %s", execResp.Error) + } + + finishResp := handleRequest(runtimeState, agentRequest{ + ID: 3, + Method: tc.finishMethod, + SessionID: sessionID, + }) + if !finishResp.Success { + t.Fatalf("%s failed: %s", tc.finishMethod, finishResp.Error) + } + closeResp := handleRequest(runtimeState, agentRequest{ + ID: 4, + Method: agentMethodCloseSession, + SessionID: sessionID, + }) + if !closeResp.Success { + t.Fatalf("closeSession failed: %s", closeResp.Error) + } + if fake.transaction == nil || !fake.transaction.closed { + t.Fatal("expected managed transaction session to close") + } + if fake.transaction.commitCalls != tc.wantCommits || fake.transaction.rollbackCalls != tc.wantRollbacks { + t.Fatalf( + "unexpected finish calls: commit=%d rollback=%d", + fake.transaction.commitCalls, + fake.transaction.rollbackCalls, + ) + } + }) + } +} + func TestHandleStreamRequest_UsesSessionStreamerAndWritesChunks(t *testing.T) { old := agentDriverType originalAsync := runAgentMemoryTrimAsync diff --git a/internal/db/dameng_impl.go b/internal/db/dameng_impl.go index a4db0ab4..37203737 100644 --- a/internal/db/dameng_impl.go +++ b/internal/db/dameng_impl.go @@ -26,6 +26,8 @@ type DamengDB struct { forwarder *ssh.LocalForwarder // Store SSH tunnel forwarder } +var _ TransactionExecerProvider = (*DamengDB)(nil) + func (d *DamengDB) getDSN(config connection.ConnectionConfig) string { // dm://user:password@host:port?schema=... // or dm://user:password@host:port @@ -223,6 +225,23 @@ func (d *DamengDB) Exec(query string) (int64, error) { return res.RowsAffected() } +// OpenTransactionExecer starts a driver-backed transaction that can remain +// open across SQL editor RPCs until an explicit commit or rollback. +func (d *DamengDB) OpenTransactionExecer(ctx context.Context) (TransactionExecer, error) { + if d.conn == nil { + return nil, fmt.Errorf("连接未打开") + } + if err := ctx.Err(); err != nil { + return nil, err + } + // Do not bind the transaction to ctx: the editor finishes it in a later RPC. + tx, err := d.conn.Begin() + if err != nil { + return nil, err + } + return NewSQLTxStatementExecer(tx), nil +} + func (d *DamengDB) GetDatabases() ([]string, error) { // 达梦在本项目中将 schema/owner 作为“数据库”展示口径。 // 先查当前 schema / 当前用户,再聚合可见用户与 owner,避免权限受限时返回空列表。 diff --git a/internal/db/dameng_transaction_test.go b/internal/db/dameng_transaction_test.go new file mode 100644 index 00000000..574d610f --- /dev/null +++ b/internal/db/dameng_transaction_test.go @@ -0,0 +1,139 @@ +//go:build gonavi_full_drivers || gonavi_dameng_driver + +package db + +import ( + "context" + "database/sql" + "database/sql/driver" + "errors" + "reflect" + "sync" + "testing" +) + +type damengTransactionRecordingState struct { + mu sync.Mutex + beginCalls int + commitCalls int + rollbackCalls int + execQueries []string +} + +type damengTransactionConnector struct { + state *damengTransactionRecordingState +} + +func (c *damengTransactionConnector) Connect(context.Context) (driver.Conn, error) { + return &damengTransactionConn{state: c.state}, nil +} + +func (c *damengTransactionConnector) Driver() driver.Driver { + return damengTransactionDriver{} +} + +type damengTransactionDriver struct{} + +func (damengTransactionDriver) Open(string) (driver.Conn, error) { + return nil, errors.New("use connector") +} + +type damengTransactionConn struct { + state *damengTransactionRecordingState +} + +func (c *damengTransactionConn) Prepare(string) (driver.Stmt, error) { + return nil, errors.New("prepare is not supported") +} + +func (c *damengTransactionConn) Close() error { return nil } + +func (c *damengTransactionConn) Begin() (driver.Tx, error) { + c.state.mu.Lock() + c.state.beginCalls++ + c.state.mu.Unlock() + return &damengTransactionTx{state: c.state}, nil +} + +func (c *damengTransactionConn) ExecContext(_ context.Context, query string, _ []driver.NamedValue) (driver.Result, error) { + c.state.mu.Lock() + c.state.execQueries = append(c.state.execQueries, query) + c.state.mu.Unlock() + return driver.RowsAffected(1), nil +} + +type damengTransactionTx struct { + state *damengTransactionRecordingState +} + +func (tx *damengTransactionTx) Commit() error { + tx.state.mu.Lock() + tx.state.commitCalls++ + tx.state.mu.Unlock() + return nil +} + +func (tx *damengTransactionTx) Rollback() error { + tx.state.mu.Lock() + tx.state.rollbackCalls++ + tx.state.mu.Unlock() + return nil +} + +func (s *damengTransactionRecordingState) snapshot() (int, int, int, []string) { + s.mu.Lock() + defer s.mu.Unlock() + return s.beginCalls, s.commitCalls, s.rollbackCalls, append([]string(nil), s.execQueries...) +} + +func TestDamengOpenTransactionExecerUsesDriverTransaction(t *testing.T) { + for _, tc := range []struct { + name string + finish func(TransactionExecer) error + wantCommits int + wantRollbacks int + }{ + {name: "commit", finish: func(tx TransactionExecer) error { return tx.Commit() }, wantCommits: 1}, + {name: "rollback", finish: func(tx TransactionExecer) error { return tx.Rollback() }, wantRollbacks: 1}, + {name: "close", finish: func(tx TransactionExecer) error { return tx.Close() }, wantRollbacks: 1}, + } { + t.Run(tc.name, func(t *testing.T) { + state := &damengTransactionRecordingState{} + dbConn := sql.OpenDB(&damengTransactionConnector{state: state}) + t.Cleanup(func() { _ = dbConn.Close() }) + + damengDB := &DamengDB{conn: dbConn} + openCtx, cancel := context.WithCancel(context.Background()) + tx, err := damengDB.OpenTransactionExecer(openCtx) + if err != nil { + cancel() + t.Fatalf("OpenTransactionExecer returned error: %v", err) + } + cancel() + + stmt := "UPDATE users SET name = 'new' WHERE id = 1" + if _, err := tx.ExecContext(context.Background(), stmt); err != nil { + t.Fatalf("DML after open context cancellation returned error: %v", err) + } + if err := tc.finish(tx); err != nil { + t.Fatalf("finish transaction returned error: %v", err) + } + if err := tx.Close(); err != nil { + t.Fatalf("Close returned error: %v", err) + } + + beginCalls, commitCalls, rollbackCalls, execQueries := state.snapshot() + if beginCalls != 1 || commitCalls != tc.wantCommits || rollbackCalls != tc.wantRollbacks { + t.Fatalf( + "unexpected transaction calls: begin=%d commit=%d rollback=%d", + beginCalls, + commitCalls, + rollbackCalls, + ) + } + if !reflect.DeepEqual(execQueries, []string{stmt}) { + t.Fatalf("expected only DML to reach the driver, got %#v", execQueries) + } + }) + } +} diff --git a/internal/db/database_optional_factories_full.go b/internal/db/database_optional_factories_full.go index 05338b21..ba8ec547 100644 --- a/internal/db/database_optional_factories_full.go +++ b/internal/db/database_optional_factories_full.go @@ -11,7 +11,7 @@ func registerOptionalDatabaseFactories() { registerDatabaseFactory(newOptionalDriverAgentDatabase("sqlserver"), "sqlserver") registerDatabaseFactory(newOptionalDriverAgentDatabase("sqlite"), "sqlite") registerDatabaseFactory(newOptionalDriverAgentDatabase("duckdb"), "duckdb") - registerDatabaseFactory(newOptionalDriverAgentDatabase("dameng"), "dameng") + registerDatabaseFactory(newOptionalDriverAgentTransactionalDatabase("dameng"), "dameng") registerDatabaseFactory(newOptionalDriverAgentDatabase("kingbase"), "kingbase") registerDatabaseFactory(newOptionalDriverAgentDatabase("highgo"), "highgo") registerDatabaseFactory(newOptionalDriverAgentDatabase("vastbase"), "vastbase") diff --git a/internal/db/database_optional_factories_lite.go b/internal/db/database_optional_factories_lite.go index 9db6ec86..4cfaea78 100644 --- a/internal/db/database_optional_factories_lite.go +++ b/internal/db/database_optional_factories_lite.go @@ -11,7 +11,7 @@ func registerOptionalDatabaseFactories() { registerDatabaseFactory(newOptionalDriverAgentDatabase("sqlserver"), "sqlserver") registerDatabaseFactory(newOptionalDriverAgentDatabase("sqlite"), "sqlite") registerDatabaseFactory(newOptionalDriverAgentDatabase("duckdb"), "duckdb") - registerDatabaseFactory(newOptionalDriverAgentDatabase("dameng"), "dameng") + registerDatabaseFactory(newOptionalDriverAgentTransactionalDatabase("dameng"), "dameng") registerDatabaseFactory(newOptionalDriverAgentDatabase("kingbase"), "kingbase") registerDatabaseFactory(newOptionalDriverAgentDatabase("highgo"), "highgo") registerDatabaseFactory(newOptionalDriverAgentDatabase("vastbase"), "vastbase") diff --git a/internal/db/driver_agent_revisions_gen.go b/internal/db/driver_agent_revisions_gen.go index 85468ec0..6590b8ca 100644 --- a/internal/db/driver_agent_revisions_gen.go +++ b/internal/db/driver_agent_revisions_gen.go @@ -4,26 +4,26 @@ package db func init() { optionalDriverAgentRevisions = map[string]string{ - "mariadb": "src-cc133d2524ceb634", - "oceanbase": "src-ac17327184366ff0", - "diros": "src-7d4fe439271d0c56", - "starrocks": "src-ce9ee22641a32f46", - "sphinx": "src-08f5ae54efb3d9df", - "sqlserver": "src-6c0e98d6d8ba439d", - "sqlite": "src-96dfa25b3042b2d5", - "duckdb": "src-8804eb2cdbc89433", - "dameng": "src-016e77082aea6718", - "kingbase": "src-17728b2ebda94dc9", - "highgo": "src-da2e8a9d2e661d3b", - "vastbase": "src-da186ac367206c16", - "opengauss": "src-54dc852e4c502947", - "gaussdb": "src-3bbbffc6991dc8ae", - "iris": "src-e798713e492e9a09", - "mongodb": "src-2610395b35c2e708", - "tdengine": "src-779b9b537f08856f", - "iotdb": "src-7edea4aba8d4869e", - "clickhouse": "src-d4150c3fb3d1313a", - "elasticsearch": "src-3dc1697786483347", - "trino": "src-ba947f211ce7b19f", + "mariadb": "src-237be74eb89cf64e", + "oceanbase": "src-e103b5465c59286a", + "diros": "src-cac1fecdcb4f3020", + "starrocks": "src-a649837c20f6b0fd", + "sphinx": "src-885a99f6ad078ecb", + "sqlserver": "src-fabd1f87aba17942", + "sqlite": "src-c75f72581f6d7ae1", + "duckdb": "src-35a3b88138f7073f", + "dameng": "src-7c3f37cbda7974f6", + "kingbase": "src-c161f698a520730e", + "highgo": "src-01b527826acd35aa", + "vastbase": "src-af5bdbd20f3ce3d2", + "opengauss": "src-eb07258bd418ae5c", + "gaussdb": "src-bd73a0616d917197", + "iris": "src-62d1bef85c82f476", + "mongodb": "src-95f55020c6481079", + "tdengine": "src-817b3f655526b25d", + "iotdb": "src-e2a3161c7c81d21c", + "clickhouse": "src-2a5d2a885efb8e99", + "elasticsearch": "src-0f9ee301ddf73d4f", + "trino": "src-35f2d68ea9b9db64", } } diff --git a/internal/db/optional_driver_agent_impl.go b/internal/db/optional_driver_agent_impl.go index 2fb2b081..77198de0 100644 --- a/internal/db/optional_driver_agent_impl.go +++ b/internal/db/optional_driver_agent_impl.go @@ -22,27 +22,30 @@ import ( ) const ( - optionalAgentMethodConnect = "connect" - optionalAgentMethodClose = "close" - optionalAgentMethodMetadata = "metadata" - optionalAgentMethodPing = "ping" - optionalAgentMethodOpenSession = "openSession" - optionalAgentMethodCloseSession = "closeSession" - optionalAgentMethodQuery = "query" - optionalAgentMethodQueryMulti = "queryMulti" - optionalAgentMethodStreamQuery = "streamQuery" - optionalAgentMethodExec = "exec" - optionalAgentMethodGetDatabases = "getDatabases" - optionalAgentMethodGetTables = "getTables" - optionalAgentMethodGetCreateStmt = "getCreateStatement" - optionalAgentMethodGetColumns = "getColumns" - optionalAgentMethodGetAllColumns = "getAllColumns" - optionalAgentMethodGetIndexes = "getIndexes" - optionalAgentMethodGetForeignKeys = "getForeignKeys" - optionalAgentMethodGetTriggers = "getTriggers" - optionalAgentMethodApplyChanges = "applyChanges" - optionalAgentDefaultScannerMaxBytes = 8 << 20 - optionalAgentMetadataProbeTimeout = 5 * time.Second + optionalAgentMethodConnect = "connect" + optionalAgentMethodClose = "close" + optionalAgentMethodMetadata = "metadata" + optionalAgentMethodPing = "ping" + optionalAgentMethodOpenSession = "openSession" + optionalAgentMethodCloseSession = "closeSession" + optionalAgentMethodOpenTransaction = "openTransaction" + optionalAgentMethodCommitTransaction = "commitTransaction" + optionalAgentMethodRollbackTransaction = "rollbackTransaction" + optionalAgentMethodQuery = "query" + optionalAgentMethodQueryMulti = "queryMulti" + optionalAgentMethodStreamQuery = "streamQuery" + optionalAgentMethodExec = "exec" + optionalAgentMethodGetDatabases = "getDatabases" + optionalAgentMethodGetTables = "getTables" + optionalAgentMethodGetCreateStmt = "getCreateStatement" + optionalAgentMethodGetColumns = "getColumns" + optionalAgentMethodGetAllColumns = "getAllColumns" + optionalAgentMethodGetIndexes = "getIndexes" + optionalAgentMethodGetForeignKeys = "getForeignKeys" + optionalAgentMethodGetTriggers = "getTriggers" + optionalAgentMethodApplyChanges = "applyChanges" + optionalAgentDefaultScannerMaxBytes = 8 << 20 + optionalAgentMetadataProbeTimeout = 5 * time.Second // callStreamQueryGCInterval 控制 callStreamQuery 每接收多少行 driver-agent 数据触发一次 runtime.GC。 // // 该路径不走 sql.Rows(scan_rows.go 的周期 GC 覆盖不到),但每个 chunk 解码 @@ -433,6 +436,10 @@ type OptionalDriverAgentDB struct { client *optionalDriverAgentClient } +type optionalDriverAgentTransactionalDB struct { + *OptionalDriverAgentDB +} + type optionalDriverAgentSession struct { client *optionalDriverAgentClient driver string @@ -441,6 +448,15 @@ type optionalDriverAgentSession struct { closed bool } +type optionalDriverAgentTransaction struct { + *optionalDriverAgentSession + finishMu sync.Mutex + finished bool +} + +var _ TransactionExecerProvider = (*optionalDriverAgentTransactionalDB)(nil) +var _ TransactionExecer = (*optionalDriverAgentTransaction)(nil) + func newOptionalDriverAgentDatabase(driverType string) databaseFactory { normalized := normalizeRuntimeDriverType(driverType) return func() Database { @@ -448,6 +464,15 @@ func newOptionalDriverAgentDatabase(driverType string) databaseFactory { } } +func newOptionalDriverAgentTransactionalDatabase(driverType string) databaseFactory { + normalized := normalizeRuntimeDriverType(driverType) + return func() Database { + return &optionalDriverAgentTransactionalDB{ + OptionalDriverAgentDB: &OptionalDriverAgentDB{driverType: normalized}, + } + } +} + func (d *OptionalDriverAgentDB) Connect(config connection.ConnectionConfig) error { if d.client != nil { _ = d.client.close() @@ -686,6 +711,64 @@ func (d *OptionalDriverAgentDB) OpenSessionExecer(ctx context.Context) (Statemen }, nil } +func (d *optionalDriverAgentTransactionalDB) OpenTransactionExecer(ctx context.Context) (TransactionExecer, error) { + if d == nil || d.OptionalDriverAgentDB == nil { + return nil, fmt.Errorf("连接未打开") + } + if err := ctx.Err(); err != nil { + return nil, err + } + client, err := d.requireClient() + if err != nil { + return nil, err + } + var sessionID string + if err := client.call(optionalAgentRequest{ + Method: optionalAgentMethodOpenTransaction, + TimeoutMs: timeoutMsFromContext(ctx), + }, &sessionID, nil, nil, nil); err != nil { + return nil, err + } + sessionID = strings.TrimSpace(sessionID) + if sessionID == "" { + return nil, fmt.Errorf("%s 驱动代理未返回事务 ID", driverDisplayName(d.driverType)) + } + return &optionalDriverAgentTransaction{ + optionalDriverAgentSession: &optionalDriverAgentSession{ + client: client, + driver: d.driverType, + sessionID: sessionID, + }, + }, nil +} + +func (t *optionalDriverAgentTransaction) Commit() error { + return t.finish(optionalAgentMethodCommitTransaction) +} + +func (t *optionalDriverAgentTransaction) Rollback() error { + return t.finish(optionalAgentMethodRollbackTransaction) +} + +func (t *optionalDriverAgentTransaction) finish(method string) error { + if t == nil || t.optionalDriverAgentSession == nil { + return nil + } + t.finishMu.Lock() + defer t.finishMu.Unlock() + if t.finished { + return nil + } + if err := t.ensureOpen(); err != nil { + return err + } + t.finished = true + return t.client.call(optionalAgentRequest{ + Method: method, + SessionID: t.sessionID, + }, nil, nil, nil, nil) +} + func (s *optionalDriverAgentSession) Query(query string) ([]map[string]interface{}, []string, error) { return s.QueryContext(context.Background(), query) } diff --git a/internal/db/optional_driver_agent_impl_test.go b/internal/db/optional_driver_agent_impl_test.go index d4a5b04c..8b381466 100644 --- a/internal/db/optional_driver_agent_impl_test.go +++ b/internal/db/optional_driver_agent_impl_test.go @@ -3,6 +3,7 @@ package db import ( "bufio" "bytes" + "context" "strings" "testing" @@ -242,3 +243,80 @@ func TestOptionalDriverAgentDBQueryMultiWithMessagesParsesResultSets(t *testing. t.Fatalf("请求未使用 queryMulti 方法: %s", stdin.String()) } } + +func TestDamengOptionalDriverAgentSupportsManagedTransactions(t *testing.T) { + damengDB, err := NewDatabase("dameng") + if err != nil { + t.Fatalf("create Dameng optional driver database: %v", err) + } + if _, ok := damengDB.(TransactionExecerProvider); !ok { + t.Fatal("expected Dameng optional driver database to expose managed transactions") + } + + for _, dbType := range []string{"sqlserver", "kingbase"} { + dbInst, err := NewDatabase(dbType) + if err != nil { + t.Fatalf("create %s optional driver database: %v", dbType, err) + } + if _, ok := dbInst.(TransactionExecerProvider); ok { + t.Fatalf("expected %s to keep using its existing session transaction path", dbType) + } + } +} + +func TestOptionalDriverAgentTransactionUsesTransactionRPC(t *testing.T) { + for _, tc := range []struct { + name string + finishMethod string + finish func(TransactionExecer) error + }{ + {name: "commit", finishMethod: optionalAgentMethodCommitTransaction, finish: func(tx TransactionExecer) error { return tx.Commit() }}, + {name: "rollback", finishMethod: optionalAgentMethodRollbackTransaction, finish: func(tx TransactionExecer) error { return tx.Rollback() }}, + } { + t.Run(tc.name, func(t *testing.T) { + var stdin optionalAgentTestWriteCloser + stdout := strings.Join([]string{ + `{"id":1,"success":true,"data":"transaction-1"}`, + `{"id":2,"success":true,"rowsAffected":1}`, + `{"id":3,"success":true}`, + `{"id":4,"success":true}`, + }, "\n") + "\n" + dbInst := &optionalDriverAgentTransactionalDB{ + OptionalDriverAgentDB: &OptionalDriverAgentDB{ + driverType: "dameng", + client: &optionalDriverAgentClient{ + stdin: &stdin, + reader: bufio.NewReader(strings.NewReader(stdout)), + driver: "dameng", + }, + }, + } + + tx, err := dbInst.OpenTransactionExecer(context.Background()) + if err != nil { + t.Fatalf("OpenTransactionExecer returned error: %v", err) + } + if _, err := tx.ExecContext(context.Background(), "UPDATE t SET v = 1"); err != nil { + t.Fatalf("ExecContext returned error: %v", err) + } + if err := tc.finish(tx); err != nil { + t.Fatalf("finish transaction returned error: %v", err) + } + if err := tx.Close(); err != nil { + t.Fatalf("Close returned error: %v", err) + } + + requests := stdin.String() + for _, fragment := range []string{ + `"method":"openTransaction"`, + `"method":"exec","sessionId":"transaction-1"`, + `"method":"` + tc.finishMethod + `","sessionId":"transaction-1"`, + `"method":"closeSession","sessionId":"transaction-1"`, + } { + if !strings.Contains(requests, fragment) { + t.Fatalf("expected request fragment %q, got %s", fragment, requests) + } + } + }) + } +}