This is an automated email from the ASF dual-hosted git repository.
thunguo pushed a commit to branch master
in repository https://gitbox.apache.org/repos/asf/incubator-seata-go.git
The following commit(s) were added to refs/heads/master by this push:
new 84c991ff Feat:support postgres in xa mode (#1117)
84c991ff is described below
commit 84c991ff94a2dd8e1d1dcaaf4e8a8cc85d24daca
Author: ssshr-66 <[email protected]>
AuthorDate: Sat Jun 6 11:20:02 2026 +0800
Feat:support postgres in xa mode (#1117)
* upd
* upd
---------
Co-authored-by: ThunGuo <[email protected]>
---
changes/dev.md | 1 +
changes/dev_zh.md | 1 +
docs/quickstart.md | 15 ++-
docs/quickstart_zh.md | 15 ++-
pkg/datasource/sql/connector.go | 35 ++++-
pkg/datasource/sql/driver.go | 81 ++++++++++--
pkg/datasource/sql/postgres_driver_test.go | 109 ++++++++++++++++
pkg/datasource/sql/xa/postgres_xa_connection.go | 145 +++++++++++++++++++--
.../sql/xa/postgres_xa_connection_test.go | 132 +++++++++++++++++++
9 files changed, 506 insertions(+), 28 deletions(-)
diff --git a/changes/dev.md b/changes/dev.md
index cfec2483..b45990ce 100755
--- a/changes/dev.md
+++ b/changes/dev.md
@@ -26,6 +26,7 @@
### feature:
- [[#123](https://github.com/apache/incubator-seata-go/pull/123)] add two
phase and dubbo
+ - support PostgreSQL XA via pgx driver
### bugfix:
diff --git a/changes/dev_zh.md b/changes/dev_zh.md
index 689aefaf..623e53c0 100644
--- a/changes/dev_zh.md
+++ b/changes/dev_zh.md
@@ -27,6 +27,7 @@ Seata-go 是一款开源的分布式事务解决方案,提供高性能和简
### feature:
- [[#123](https://github.com/apache/incubator-seata-go/pull/123)]
添加二阶段事务接口,以及dubbo集成
+- 支持基于 pgx 驱动的 PostgreSQL XA
### bugfix:
diff --git a/docs/quickstart.md b/docs/quickstart.md
index 8ade97fc..7717eb9c 100644
--- a/docs/quickstart.md
+++ b/docs/quickstart.md
@@ -273,7 +273,9 @@ If you need fence mode, set `seata.tcc.fence.enable: true`
in `seatago.yml` and
### XA Example
-The current repository supports MySQL XA. The transaction entrypoint is the
same as AT; you only need to switch the driver to `seata-xa-mysql`:
+The transaction entrypoint is the same as AT.
+
+For MySQL XA, switch the driver to `seata-xa-mysql`:
```go
db, err := sql.Open(
@@ -282,4 +284,15 @@ db, err := sql.Open(
)
```
+For PostgreSQL XA, use the pgx-based driver `seata-xa-postgres`:
+
+```go
+db, err := sql.Open(
+ "seata-xa-postgres",
+
"postgres://postgres:[email protected]:5432/seata_demo?sslmode=disable",
+)
+```
+
+> PostgreSQL XA relies on prepared transactions. Set
`max_prepared_transactions > 0` on the PostgreSQL server before using this mode.
+
> Full example: [XA
> Example](https://github.com/apache/incubator-seata-go-samples/tree/main/xa/basic).
diff --git a/docs/quickstart_zh.md b/docs/quickstart_zh.md
index ea0136fb..3729c810 100644
--- a/docs/quickstart_zh.md
+++ b/docs/quickstart_zh.md
@@ -273,7 +273,9 @@ return tm.WithGlobalTx(ctx, &tm.GtxConfig{Name:
"inventory-tcc"}, func(ctx conte
### XA 模式示例
-当前仓库已支持 MySQL XA,事务入口与 AT 相同,只需要将驱动切换为 `seata-xa-mysql`:
+XA 模式的事务入口与 AT 相同。
+
+如果使用 MySQL XA,只需要将驱动切换为 `seata-xa-mysql`:
```go
db, err := sql.Open(
@@ -282,4 +284,15 @@ db, err := sql.Open(
)
```
+如果使用 PostgreSQL XA,可使用基于 pgx 的驱动 `seata-xa-postgres`:
+
+```go
+db, err := sql.Open(
+ "seata-xa-postgres",
+
"postgres://postgres:[email protected]:5432/seata_demo?sslmode=disable",
+)
+```
+
+> PostgreSQL XA 依赖 prepared transaction,使用前需要在 PostgreSQL 服务端开启
`max_prepared_transactions > 0`。
+
> 完整示例可参考:[XA
> 模式示例](https://github.com/apache/incubator-seata-go-samples/tree/main/xa/basic)。
diff --git a/pkg/datasource/sql/connector.go b/pkg/datasource/sql/connector.go
index 7474fc21..90a8a83d 100644
--- a/pkg/datasource/sql/connector.go
+++ b/pkg/datasource/sql/connector.go
@@ -20,12 +20,15 @@ package sql
import (
"context"
"database/sql/driver"
+ "sync"
"seata.apache.org/seata-go/v2/pkg/datasource/sql/types"
+ "seata.apache.org/seata-go/v2/pkg/protocol/branch"
)
type seataATConnector struct {
*seataConnector
+ transType types.TransactionMode
}
func (c *seataATConnector) Connect(ctx context.Context) (driver.Conn, error) {
@@ -49,6 +52,7 @@ func (c *seataATConnector) Driver() driver.Driver {
type seataXAConnector struct {
*seataConnector
+ transType types.TransactionMode
}
func (c *seataXAConnector) Connect(ctx context.Context) (driver.Conn, error) {
@@ -83,12 +87,16 @@ func (c *seataXAConnector) Driver() driver.Driver {
// If a Connector implements io.Closer, the sql package's DB.Close
// method will call Close and return error (if any).
type seataConnector struct {
- transType types.TransactionMode
- res *DBResource
- driver *seataDriver
- target driver.Connector
- dbType types.DBType
- dbName string
+ transType types.TransactionMode
+ branchType branch.BranchType
+ res *DBResource
+ driver driver.Driver
+ target driver.Connector
+ targetDriver driver.Driver
+ targetName string
+ dbType types.DBType
+ dbName string
+ once sync.Once
}
// Connect returns a connection to the database.
@@ -123,5 +131,18 @@ func (c *seataConnector) Connect(ctx context.Context)
(driver.Conn, error) {
// mainly to maintain compatibility with the Driver method
// on sql.DB.
func (c *seataConnector) Driver() driver.Driver {
- return c.driver
+ c.once.Do(func() {
+ if c.targetDriver != nil {
+ c.driver = c.targetDriver
+ return
+ }
+ c.driver = c.target.Driver()
+ })
+
+ return &seataDriver{
+ branchType: c.branchType,
+ transType: c.transType,
+ target: c.driver,
+ targetName: c.targetName,
+ }
}
diff --git a/pkg/datasource/sql/driver.go b/pkg/datasource/sql/driver.go
index 7f70ad68..de796c53 100644
--- a/pkg/datasource/sql/driver.go
+++ b/pkg/datasource/sql/driver.go
@@ -29,7 +29,7 @@ import (
"github.com/go-sql-driver/mysql"
"github.com/jackc/pgx/v5"
- pgxstdlib "github.com/jackc/pgx/v5/stdlib"
+ "github.com/jackc/pgx/v5/stdlib"
"seata.apache.org/seata-go/v2/pkg/datasource/sql/datasource"
mysql2
"seata.apache.org/seata-go/v2/pkg/datasource/sql/datasource/mysql"
@@ -47,6 +47,8 @@ const (
SeataATPostgresDriver = "seata-at-postgres"
// SeataXAMySQLDriver MySQL driver for XA mode
SeataXAMySQLDriver = "seata-xa-mysql"
+ // SeataXAPostgresDriver PostgreSQL driver for XA mode
+ SeataXAPostgresDriver = "seata-xa-postgres"
)
type driverDescriptor struct {
@@ -67,7 +69,7 @@ var (
}
postgresDriverDescriptor = driverDescriptor{
dbType: types.DBTypePostgreSQL,
- target: pgxstdlib.GetDefaultDriver(),
+ target: stdlib.GetDefaultDriver(),
parseDBName: parsePostgresDBName,
newTableMetaCache: func(db *sql.DB, dbName string)
datasource.TableMetaCache {
return postgres2.NewTableMetaInstance(db, dbName)
@@ -80,6 +82,8 @@ func initDriver() {
seataDriver: &seataDriver{
branchType: branch.BranchTypeAT,
transType: types.ATMode,
+ target: mysql.MySQLDriver{},
+ targetName: "mysql",
descriptor: mySQLDriverDescriptor,
},
})
@@ -97,6 +101,18 @@ func initDriver() {
branchType: branch.BranchTypeXA,
transType: types.XAMode,
descriptor: mySQLDriverDescriptor,
+ target: mysql.MySQLDriver{},
+ targetName: "mysql",
+ },
+ })
+
+ sql.Register(SeataXAPostgresDriver, &seataXADriver{
+ seataDriver: &seataDriver{
+ branchType: branch.BranchTypeXA,
+ transType: types.XAMode,
+ descriptor: postgresDriverDescriptor,
+ target: stdlib.GetDefaultDriver(),
+ targetName: "pgx",
},
})
}
@@ -112,6 +128,7 @@ func (d *seataATDriver) OpenConnector(name string) (c
driver.Connector, err erro
}
_connector, _ := connector.(*seataConnector)
+ _connector.transType = types.ATMode
return &seataATConnector{
seataConnector: _connector,
@@ -129,6 +146,7 @@ func (d *seataXADriver) OpenConnector(name string) (c
driver.Connector, err erro
}
_connector, _ := connector.(*seataConnector)
+ _connector.transType = types.XAMode
return &seataXAConnector{
seataConnector: _connector,
@@ -139,6 +157,8 @@ type seataDriver struct {
branchType branch.BranchType
transType types.TransactionMode
descriptor driverDescriptor
+ target driver.Driver
+ targetName string
}
// Open never be called, because seataDriver implemented dri.DriverContext
interface.
@@ -174,6 +194,11 @@ func (d *seataDriver) OpenConnector(name string) (c
driver.Connector, err error)
func (d *seataDriver) getOpenConnectorProxy(connector driver.Connector, dbType
types.DBType,
db *sql.DB, dataSourceName string) (driver.Connector, error) {
+ meta, err := parseConnectorMetadata(dataSourceName, dbType)
+ if err != nil {
+ return nil, err
+ }
+
dbName, err := d.descriptor.parseDBName(dataSourceName)
if err != nil {
return nil, fmt.Errorf("parse db name: %w", err)
@@ -184,6 +209,7 @@ func (d *seataDriver) getOpenConnectorProxy(connector
driver.Connector, dbType t
withBranchType(d.branchType),
withDBType(dbType),
withDBName(dbName),
+ withDBName(meta.dbName),
withConnector(connector),
}
res, err := newResource(options...)
@@ -191,18 +217,30 @@ func (d *seataDriver) getOpenConnectorProxy(connector
driver.Connector, dbType t
log.Errorf("create new resource: %v", err)
return nil, err
}
+
+ if dbType == types.DBTypeMySQL {
+ cfg, err := mysql.ParseDSN(dataSourceName)
+ if err != nil {
+ return nil, fmt.Errorf("parse mysql dsn: %w", err)
+ }
+ datasource.RegisterTableCache(types.DBTypeMySQL,
mysql2.NewTableMetaInstance(db, cfg))
+ }
+
datasource.RegisterTableCache(dbType,
d.descriptor.newTableMetaCache(db, dbName))
if err =
datasource.GetDataSourceManager(d.branchType).RegisterResource(res); err != nil
{
log.Errorf("register resource: %v", err)
return nil, err
}
return &seataConnector{
- transType: d.transType,
- res: res,
- driver: d,
- target: connector,
- dbType: dbType,
- dbName: dbName,
+ transType: d.transType,
+ branchType: d.branchType,
+ res: res,
+ driver: d,
+ target: connector,
+ targetDriver: d.target,
+ targetName: d.targetName,
+ dbType: dbType,
+ dbName: dbName,
}, nil
}
@@ -222,6 +260,33 @@ func parsePostgresDBName(dsn string) (string, error) {
return cfg.Database, nil
}
+func (d *seataDriver) getTargetDriverName() string {
+ return d.targetName
+}
+
+type connectorMetadata struct {
+ dbName string
+}
+
+func parseConnectorMetadata(dataSourceName string, dbType types.DBType)
(*connectorMetadata, error) {
+ switch dbType {
+ case types.DBTypeMySQL:
+ cfg, err := mysql.ParseDSN(dataSourceName)
+ if err != nil {
+ return nil, fmt.Errorf("parse mysql dsn: %w", err)
+ }
+ return &connectorMetadata{dbName: cfg.DBName}, nil
+ case types.DBTypePostgreSQL:
+ cfg, err := pgx.ParseConfig(dataSourceName)
+ if err != nil {
+ return nil, fmt.Errorf("parse postgres dsn: %w", err)
+ }
+ return &connectorMetadata{dbName: cfg.Database}, nil
+ default:
+ return nil, fmt.Errorf("unsupported connector metadata for db
type %s", dbType.String())
+ }
+}
+
type dsnConnector struct {
dsn string
driver driver.Driver
diff --git a/pkg/datasource/sql/postgres_driver_test.go
b/pkg/datasource/sql/postgres_driver_test.go
new file mode 100644
index 00000000..87bfb7e4
--- /dev/null
+++ b/pkg/datasource/sql/postgres_driver_test.go
@@ -0,0 +1,109 @@
+/*
+ * Licensed to the Apache Software Foundation (ASF) under one or more
+ * contributor license agreements. See the NOTICE file distributed with
+ * this work for additional information regarding copyright ownership.
+ * The ASF licenses this file to You under the Apache License, Version 2.0
+ * (the "License"); you may not use this file except in compliance with
+ * the License. You may obtain a copy of the License at
+ *
+ * http://www.apache.org/licenses/LICENSE-2.0
+ *
+ * Unless required by applicable law or agreed to in writing, software
+ * distributed under the License is distributed on an "AS IS" BASIS,
+ * WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied.
+ * See the License for the specific language governing permissions and
+ * limitations under the License.
+ */
+
+package sql
+
+import (
+ "context"
+ "testing"
+
+ "github.com/golang/mock/gomock"
+ "github.com/stretchr/testify/assert"
+
+ "seata.apache.org/seata-go/v2/pkg/datasource/sql/mock"
+ "seata.apache.org/seata-go/v2/pkg/datasource/sql/types"
+)
+
+func TestParseConnectorMetadataPostgres(t *testing.T) {
+ meta, err :=
parseConnectorMetadata("postgres://postgres:[email protected]:5432/seata_demo?sslmode=disable",
types.DBTypePostgreSQL)
+ assert.NoError(t, err)
+ assert.Equal(t, "seata_demo", meta.dbName)
+}
+
+func TestSeataConnectorConnectPostgres(t *testing.T) {
+ ctrl := gomock.NewController(t)
+ defer ctrl.Finish()
+
+ mockConn := mock.NewMockTestDriverConn(ctrl)
+ mockConnector := mock.NewMockTestDriverConnector(ctrl)
+ mockConnector.EXPECT().Connect(gomock.Any()).Return(mockConn, nil)
+
+ connector := &seataConnector{
+ target: mockConnector,
+ res: &DBResource{
+ dbType: types.DBTypePostgreSQL,
+ },
+ dbName: "seata_demo",
+ dbType: types.DBTypePostgreSQL,
+ }
+
+ conn, err := connector.Connect(context.Background())
+ assert.NoError(t, err)
+
+ got, ok := conn.(*Conn)
+ assert.True(t, ok)
+ assert.Equal(t, "seata_demo", got.dbName)
+ assert.Equal(t, types.DBTypePostgreSQL, got.dbType)
+}
+
+func TestSeataXAConnectorConnectPostgres(t *testing.T) {
+ ctrl := gomock.NewController(t)
+ defer ctrl.Finish()
+
+ mockConn := mock.NewMockTestDriverConn(ctrl)
+ mockConnector := mock.NewMockTestDriverConnector(ctrl)
+ mockConnector.EXPECT().Connect(gomock.Any()).Return(mockConn, nil)
+
+ connector := &seataXAConnector{
+ seataConnector: &seataConnector{
+ target: mockConnector,
+ res: &DBResource{
+ dbType: types.DBTypePostgreSQL,
+ },
+ dbName: "seata_demo",
+ dbType: types.DBTypePostgreSQL,
+ },
+ }
+
+ conn, err := connector.Connect(context.Background())
+ assert.NoError(t, err)
+
+ xaConn, ok := conn.(*XAConn)
+ assert.True(t, ok)
+ assert.Equal(t, "seata_demo", xaConn.dbName)
+ assert.Equal(t, types.DBTypePostgreSQL, xaConn.dbType)
+ assert.True(t, xaConn.txCtx.TransactionMode == types.Local)
+}
+
+func TestSeataConnectorDriverPreservesTargetName(t *testing.T) {
+ ctrl := gomock.NewController(t)
+ defer ctrl.Finish()
+
+ mockDriver := mock.NewMockTestDriver(ctrl)
+ mockConnector := mock.NewMockTestDriverConnector(ctrl)
+
+ connector := &seataConnector{
+ target: mockConnector,
+ targetDriver: mockDriver,
+ targetName: "pgx",
+ transType: types.XAMode,
+ }
+
+ got, ok := connector.Driver().(*seataDriver)
+ assert.True(t, ok)
+ assert.Equal(t, "pgx", got.getTargetDriverName())
+}
diff --git a/pkg/datasource/sql/xa/postgres_xa_connection.go
b/pkg/datasource/sql/xa/postgres_xa_connection.go
index a1a0fb51..a33b5764 100644
--- a/pkg/datasource/sql/xa/postgres_xa_connection.go
+++ b/pkg/datasource/sql/xa/postgres_xa_connection.go
@@ -20,7 +20,10 @@ package xa
import (
"context"
"database/sql/driver"
+ "errors"
"fmt"
+ "io"
+ "strings"
"time"
"seata.apache.org/seata-go/v2/pkg/datasource/sql/types"
@@ -60,39 +63,155 @@ func (c *PostgresXAErrorClassifier) IsAlreadyEnded(err
error) bool {
// - Requires max_prepared_transactions > 0 in postgresql.conf.
type PostgresXAConn struct {
driver.Conn
+ tx driver.Tx
}
func (c *PostgresXAConn) Start(ctx context.Context, xid string, flags int)
error {
- log.Infof("xa branch start (postgres no-op), xid %s", xid)
+ log.Infof("xa branch start (postgres begin), xid %s", xid)
+
+ if flags != TMNoFlags {
+ return errors.New("invalid arguments")
+ }
+ if c.tx != nil {
+ return fmt.Errorf("postgres xa transaction already started, xid
%s", xid)
+ }
+
+ var (
+ tx driver.Tx
+ err error
+ )
+ if conn, ok := c.Conn.(driver.ConnBeginTx); ok {
+ tx, err = conn.BeginTx(ctx, driver.TxOptions{})
+ } else {
+ tx, err = c.Conn.Begin()
+ }
+ if err != nil {
+ log.Errorf("postgres xa branch start failed, xid %s, err %v",
xid, err)
+ return err
+ }
+
+ c.tx = tx
return nil
}
func (c *PostgresXAConn) End(ctx context.Context, xid string, flags int) error
{
log.Infof("xa branch end (postgres no-op), xid %s", xid)
- return nil
+
+ switch flags {
+ case TMSuccess, TMFail:
+ return nil
+ default:
+ return errors.New("invalid arguments")
+ }
}
func (c *PostgresXAConn) XAPrepare(ctx context.Context, xid string) error {
- // TODO: PREPARE TRANSACTION 'xid'
- return fmt.Errorf("PostgreSQL XA prepare not yet implemented")
+ log.Infof("postgres xa branch prepare, xid %s", xid)
+
+ if c.tx == nil {
+ return fmt.Errorf("postgres xa prepare requires active
transaction, xid %s", xid)
+ }
+
+ query := "PREPARE TRANSACTION " + quotePostgresXID(xid)
+ conn, _ := c.Conn.(driver.ExecerContext)
+ _, err := conn.ExecContext(ctx, query, nil)
+ if err != nil {
+ log.Errorf("postgres xa branch prepare failed, xid %s, err %v",
xid, err)
+ return err
+ }
+
+ c.tx = nil
+ return nil
}
func (c *PostgresXAConn) Commit(ctx context.Context, xid string, onePhase
bool) error {
- // TODO: COMMIT PREPARED 'xid' (onePhase=true → regular commit by upper
layer)
- return fmt.Errorf("PostgreSQL XA commit not yet implemented")
+ if onePhase {
+ if c.tx == nil {
+ return fmt.Errorf("postgres xa one-phase commit
requires active transaction, xid %s", xid)
+ }
+ if err := c.tx.Commit(); err != nil {
+ log.Errorf("postgres xa one-phase commit failed, xid
%s, err %v", xid, err)
+ return err
+ }
+ c.tx = nil
+ return nil
+ }
+
+ if c.tx != nil {
+ return fmt.Errorf("postgres xa commit requires prepared
transaction, xid %s", xid)
+ }
+
+ log.Infof("postgres xa branch commit prepared, xid %s", xid)
+
+ query := "COMMIT PREPARED " + quotePostgresXID(xid)
+ conn, _ := c.Conn.(driver.ExecerContext)
+ _, err := conn.ExecContext(ctx, query, nil)
+ if err != nil {
+ log.Errorf("postgres xa branch commit prepared failed, xid %s,
err %v", xid, err)
+ }
+ return err
}
func (c *PostgresXAConn) Rollback(ctx context.Context, xid string) error {
- // TODO: ROLLBACK PREPARED 'xid'
- return fmt.Errorf("PostgreSQL XA rollback not yet implemented")
+ if c.tx != nil {
+ log.Infof("postgres xa branch rollback active transaction, xid
%s", xid)
+ err := c.tx.Rollback()
+ if err != nil {
+ log.Errorf("postgres xa branch rollback active
transaction failed, xid %s, err %v", xid, err)
+ return err
+ }
+ c.tx = nil
+ return nil
+ }
+
+ log.Infof("postgres xa branch rollback prepared, xid %s", xid)
+
+ query := "ROLLBACK PREPARED " + quotePostgresXID(xid)
+ conn, _ := c.Conn.(driver.ExecerContext)
+ _, err := conn.ExecContext(ctx, query, nil)
+ if err != nil {
+ log.Errorf("postgres xa branch rollback prepared failed, xid
%s, err %v", xid, err)
+ }
+ return err
}
func (c *PostgresXAConn) Recover(ctx context.Context, flag int) ([]string,
error) {
- if (flag & TMStartRScan) == 0 {
+ startRscan := (flag & TMStartRScan) > 0
+ endRscan := (flag & TMEndRScan) > 0
+
+ if !startRscan && !endRscan && flag != TMNoFlags {
+ return nil, errors.New("invalid arguments")
+ }
+ if !startRscan {
return nil, nil
}
- // TODO: SELECT gid FROM pg_prepared_xacts WHERE database =
current_database()
- return nil, fmt.Errorf("PostgreSQL XA recover not yet implemented")
+
+ conn := c.Conn.(driver.QueryerContext)
+ rows, err := conn.QueryContext(ctx, "SELECT gid FROM pg_prepared_xacts
WHERE database = current_database()", nil)
+ if err != nil {
+ return nil, err
+ }
+ defer rows.Close()
+
+ xids := make([]string, 0)
+ dest := make([]driver.Value, 1)
+ for {
+ if err = rows.Next(dest); err != nil {
+ if err == io.EOF {
+ return xids, nil
+ }
+ return nil, err
+ }
+
+ switch v := dest[0].(type) {
+ case string:
+ xids = append(xids, v)
+ case []byte:
+ xids = append(xids, string(v))
+ default:
+ return nil, errors.New("the protocol of postgres
prepared transaction query is error")
+ }
+ }
}
func (c *PostgresXAConn) Forget(ctx context.Context, xid string) error {
@@ -104,3 +223,7 @@ func (c *PostgresXAConn) GetTransactionTimeout()
time.Duration { return 0 }
func (c *PostgresXAConn) IsSameRM(ctx context.Context, resource XAResource)
bool { return false }
func (c *PostgresXAConn) SetTransactionTimeout(duration time.Duration) bool {
return false }
+
+func quotePostgresXID(xid string) string {
+ return "'" + strings.ReplaceAll(xid, "'", "''") + "'"
+}
diff --git a/pkg/datasource/sql/xa/postgres_xa_connection_test.go
b/pkg/datasource/sql/xa/postgres_xa_connection_test.go
new file mode 100644
index 00000000..b640a8f3
--- /dev/null
+++ b/pkg/datasource/sql/xa/postgres_xa_connection_test.go
@@ -0,0 +1,132 @@
+/*
+ * Licensed to the Apache Software Foundation (ASF) under one or more
+ * contributor license agreements. See the NOTICE file distributed with
+ * this work for additional information regarding copyright ownership.
+ * The ASF licenses this file to You under the Apache License, Version 2.0
+ * (the "License"); you may not use this file except in compliance with
+ * the License. You may obtain a copy of the License at
+ *
+ * http://www.apache.org/licenses/LICENSE-2.0
+ *
+ * Unless required by applicable law or agreed to in writing, software
+ * distributed under the License is distributed on an "AS IS" BASIS,
+ * WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied.
+ * See the License for the specific language governing permissions and
+ * limitations under the License.
+ */
+
+package xa
+
+import (
+ "context"
+ "database/sql/driver"
+ "io"
+ "testing"
+
+ "github.com/golang/mock/gomock"
+ "github.com/stretchr/testify/assert"
+
+ "seata.apache.org/seata-go/v2/pkg/datasource/sql/mock"
+)
+
+type postgresMockRows struct {
+ idx int
+ data [][]interface{}
+}
+
+func (m *postgresMockRows) Columns() []string { return []string{"gid"} }
+
+func (m *postgresMockRows) Close() error { return nil }
+
+func (m *postgresMockRows) Next(dest []driver.Value) error {
+ if m.idx == len(m.data) {
+ return io.EOF
+ }
+
+ for i := 0; i < len(dest) && i < len(m.data[m.idx]); i++ {
+ dest[i] = m.data[m.idx][i]
+ }
+ m.idx++
+ return nil
+}
+
+func TestPostgresXAConn_StartAndPrepare(t *testing.T) {
+ ctrl := gomock.NewController(t)
+ defer ctrl.Finish()
+
+ mockConn := mock.NewMockTestDriverConn(ctrl)
+ mockTx := mock.NewMockTestDriverTx(ctrl)
+ mockConn.EXPECT().BeginTx(gomock.Any(), gomock.Any()).Return(mockTx,
nil)
+ mockConn.EXPECT().ExecContext(gomock.Any(), "PREPARE TRANSACTION
'xid'", gomock.Any()).Return(&driver.ResultNoRows, nil)
+
+ conn := &PostgresXAConn{Conn: mockConn}
+ assert.NoError(t, conn.Start(context.Background(), "xid", TMNoFlags))
+ assert.NoError(t, conn.XAPrepare(context.Background(), "xid"))
+ assert.Nil(t, conn.tx)
+}
+
+func TestPostgresXAConn_CommitPrepared(t *testing.T) {
+ ctrl := gomock.NewController(t)
+ defer ctrl.Finish()
+
+ mockConn := mock.NewMockTestDriverConn(ctrl)
+ mockConn.EXPECT().ExecContext(gomock.Any(), "COMMIT PREPARED 'xid'",
gomock.Any()).Return(&driver.ResultNoRows, nil)
+
+ conn := &PostgresXAConn{Conn: mockConn}
+ assert.NoError(t, conn.Commit(context.Background(), "xid", false))
+}
+
+func TestPostgresXAConn_OnePhaseCommit(t *testing.T) {
+ ctrl := gomock.NewController(t)
+ defer ctrl.Finish()
+
+ mockConn := mock.NewMockTestDriverConn(ctrl)
+ mockTx := mock.NewMockTestDriverTx(ctrl)
+ mockConn.EXPECT().BeginTx(gomock.Any(), gomock.Any()).Return(mockTx,
nil)
+ mockTx.EXPECT().Commit().Return(nil)
+
+ conn := &PostgresXAConn{Conn: mockConn}
+ assert.NoError(t, conn.Start(context.Background(), "xid", TMNoFlags))
+ assert.NoError(t, conn.Commit(context.Background(), "xid", true))
+ assert.Nil(t, conn.tx)
+}
+
+func TestPostgresXAConn_RollbackActiveTransaction(t *testing.T) {
+ ctrl := gomock.NewController(t)
+ defer ctrl.Finish()
+
+ mockConn := mock.NewMockTestDriverConn(ctrl)
+ mockTx := mock.NewMockTestDriverTx(ctrl)
+ mockConn.EXPECT().BeginTx(gomock.Any(), gomock.Any()).Return(mockTx,
nil)
+ mockTx.EXPECT().Rollback().Return(nil)
+
+ conn := &PostgresXAConn{Conn: mockConn}
+ assert.NoError(t, conn.Start(context.Background(), "xid", TMNoFlags))
+ assert.NoError(t, conn.Rollback(context.Background(), "xid"))
+ assert.Nil(t, conn.tx)
+}
+
+func TestPostgresXAConn_RollbackPrepared(t *testing.T) {
+ ctrl := gomock.NewController(t)
+ defer ctrl.Finish()
+
+ mockConn := mock.NewMockTestDriverConn(ctrl)
+ mockConn.EXPECT().ExecContext(gomock.Any(), "ROLLBACK PREPARED 'xid'",
gomock.Any()).Return(&driver.ResultNoRows, nil)
+
+ conn := &PostgresXAConn{Conn: mockConn}
+ assert.NoError(t, conn.Rollback(context.Background(), "xid"))
+}
+
+func TestPostgresXAConn_Recover(t *testing.T) {
+ ctrl := gomock.NewController(t)
+ defer ctrl.Finish()
+
+ mockConn := mock.NewMockTestDriverConn(ctrl)
+ mockConn.EXPECT().QueryContext(gomock.Any(), "SELECT gid FROM
pg_prepared_xacts WHERE database = current_database()", gomock.Any()).
+ Return(&postgresMockRows{data: [][]interface{}{{"xid"},
{"another-xid"}}}, nil)
+
+ conn := &PostgresXAConn{Conn: mockConn}
+ got, err := conn.Recover(context.Background(), TMStartRScan|TMEndRScan)
+ assert.NoError(t, err)
+ assert.Equal(t, []string{"xid", "another-xid"}, got)
+}
---------------------------------------------------------------------
To unsubscribe, e-mail: [email protected]
For additional commands, e-mail: [email protected]