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]

Reply via email to