Skip to content

Commit a361533

Browse files
committed
Fix query attribute lifecycle and prepared statement encoding
1 parent aef8a3b commit a361533

3 files changed

Lines changed: 116 additions & 5 deletions

File tree

‎client/conn.go‎

Lines changed: 1 addition & 1 deletion
Original file line numberDiff line numberDiff line change
@@ -715,7 +715,7 @@ func (c *Conn) exec(query string) (*mysql.Result, error) {
715715
// https://dev.mysql.com/doc/dev/mysql-server/latest/page_protocol_com_query.html
716716
func (c *Conn) execSend(query string) error {
717717
var buf bytes.Buffer
718-
defer clear(c.queryAttributes)
718+
defer func() { c.queryAttributes = nil }()
719719

720720
if c.capability&mysql.CLIENT_QUERY_ATTRIBUTES > 0 {
721721
if c.includeLine >= 0 {
Lines changed: 107 additions & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -0,0 +1,107 @@
1+
package client
2+
3+
import (
4+
"bytes"
5+
"net"
6+
"testing"
7+
"time"
8+
9+
"github.com/go-mysql-org/go-mysql/mysql"
10+
"github.com/go-mysql-org/go-mysql/packet"
11+
)
12+
13+
func captureAttributeRequest(t *testing.T, c *Conn, peer *packet.Conn, send func() error) []byte {
14+
t.Helper()
15+
done := make(chan error, 1)
16+
go func() { done <- send() }()
17+
peer.ResetSequence()
18+
payload, err := peer.ReadPacket()
19+
if err != nil {
20+
t.Fatal(err)
21+
}
22+
if err := peer.WritePacket([]byte{0, 0, 0, 0, 0, 0, 0, 2, 0, 0, 0}); err != nil {
23+
t.Fatal(err)
24+
}
25+
if err := <-done; err != nil {
26+
t.Fatal(err)
27+
}
28+
return payload
29+
}
30+
31+
func attributeConnection(t *testing.T) (*Conn, *packet.Conn) {
32+
a, b := net.Pipe()
33+
_ = a.SetDeadline(time.Now().Add(5 * time.Second))
34+
_ = b.SetDeadline(time.Now().Add(5 * time.Second))
35+
t.Cleanup(func() { a.Close(); b.Close() })
36+
return &Conn{Conn: packet.NewConn(a), capability: mysql.CLIENT_QUERY_ATTRIBUTES, includeLine: -1}, packet.NewConn(b)
37+
}
38+
39+
func TestQueryAttributeAttributesConsumed(t *testing.T) {
40+
c, p := attributeConnection(t)
41+
c.SetQueryAttributes(mysql.QueryAttribute{Name: "trace", Value: "first"})
42+
first := captureAttributeRequest(t, c, p, func() error { _, err := c.Execute("SELECT 1"); return err })
43+
second := captureAttributeRequest(t, c, p, func() error { _, err := c.Execute("SELECT 2"); return err })
44+
t.Logf("first=%x second=%x", first, second)
45+
if second[1] != 0 {
46+
t.Fatalf("next query still announces %d attributes; want zero", second[1])
47+
}
48+
}
49+
50+
func TestQueryAttributeZeroParameterAttributes(t *testing.T) {
51+
c, p := attributeConnection(t)
52+
c.SetQueryAttributes(mysql.QueryAttribute{Name: "trace", Value: "marker"})
53+
s := &Stmt{conn: c}
54+
data := captureAttributeRequest(t, c, p, func() error { _, err := s.Execute(); return err })
55+
t.Logf("execute=%x", data)
56+
if !bytes.Contains(data, []byte("marker")) {
57+
t.Fatal("zero-parameter execution lost its query attribute")
58+
}
59+
}
60+
61+
func TestQueryAttributeNullParameterAttributes(t *testing.T) {
62+
c, p := attributeConnection(t)
63+
c.SetQueryAttributes(mysql.QueryAttribute{Name: "trace", Value: "marker"})
64+
s := &Stmt{conn: c}
65+
s.Params = 1
66+
data := captureAttributeRequest(t, c, p, func() error { _, err := s.Execute(nil); return err })
67+
t.Logf("execute=%x", data)
68+
if !bytes.Contains(data, []byte("marker")) {
69+
t.Fatal("NULL-parameter execution lost its query attribute")
70+
}
71+
}
72+
73+
func TestQueryAttributeCallerAttributesReusable(t *testing.T) {
74+
c, p := attributeConnection(t)
75+
attrs := []mysql.QueryAttribute{{Name: "trace", Value: "marker"}}
76+
c.SetQueryAttributes(attrs...)
77+
captureAttributeRequest(t, c, p, func() error { _, err := c.Execute("SELECT 1"); return err })
78+
if attrs[0].Name != "trace" || attrs[0].Value != "marker" {
79+
t.Fatal("sending query mutated caller attribute slice")
80+
}
81+
}
82+
83+
func TestQueryAttributeLineAttributeDoesNotAccumulate(t *testing.T) {
84+
c, p := attributeConnection(t)
85+
c.IncludeLine(0)
86+
for i := 0; i < 3; i++ {
87+
data := captureAttributeRequest(t, c, p, func() error { _, err := c.Execute("SELECT 1"); return err })
88+
if data[1] != 1 {
89+
t.Fatalf("query %d has %d attributes; want one line attribute", i, data[1])
90+
}
91+
}
92+
}
93+
94+
func TestQueryAttributeUnsupportedCapability(t *testing.T) {
95+
c, p := attributeConnection(t)
96+
c.capability = 0
97+
c.SetQueryAttributes(mysql.QueryAttribute{Name: "trace", Value: "marker"})
98+
s := &Stmt{conn: c}
99+
s.Params = 1
100+
data := captureAttributeRequest(t, c, p, func() error { _, err := s.Execute(7); return err })
101+
if data[5] != 0 {
102+
t.Fatalf("legacy execute has unsupported flags %x", data[5])
103+
}
104+
if bytes.Contains(data, []byte("marker")) {
105+
t.Fatal("attributes encoded without negotiated support")
106+
}
107+
}

‎client/stmt.go‎

Lines changed: 8 additions & 4 deletions
Original file line numberDiff line numberDiff line change
@@ -102,7 +102,7 @@ func (s *Stmt) Close() error {
102102

103103
// https://dev.mysql.com/doc/dev/mysql-server/latest/page_protocol_com_stmt_execute.html
104104
func (s *Stmt) write(args ...any) error {
105-
defer clear(s.conn.queryAttributes)
105+
defer func() { s.conn.queryAttributes = nil }()
106106
paramsNum := s.Params
107107

108108
if len(args) != paramsNum {
@@ -121,6 +121,9 @@ func (s *Stmt) write(args ...any) error {
121121
}
122122

123123
qaLen := len(s.conn.queryAttributes)
124+
if s.conn.capability&mysql.CLIENT_QUERY_ATTRIBUTES == 0 {
125+
qaLen = 0
126+
}
124127
paramTypes := make([][]byte, paramsNum+qaLen)
125128
paramFlags := make([][]byte, paramsNum+qaLen)
126129
paramValues := make([][]byte, paramsNum+qaLen)
@@ -215,12 +218,13 @@ func (s *Stmt) write(args ...any) error {
215218

216219
length += len(paramValues[i])
217220
}
218-
for i, qa := range s.conn.queryAttributes {
221+
for i, qa := range s.conn.queryAttributes[:qaLen] {
219222
tf := qa.TypeAndFlag()
220223
paramTypes[(i + paramsNum)] = []byte{tf[0]}
221224
paramFlags[i+paramsNum] = []byte{tf[1]}
222225
paramValues[i+paramsNum] = qa.ValueBytes()
223226
paramNames[i+paramsNum] = mysql.PutLengthEncodedString([]byte(qa.Name))
227+
newParamBoundFlag = 1
224228
}
225229

226230
data := utils.BytesBufferGet()
@@ -236,7 +240,7 @@ func (s *Stmt) write(args ...any) error {
236240
data.Write([]byte{byte(s.ID), byte(s.ID >> 8), byte(s.ID >> 16), byte(s.ID >> 24)})
237241

238242
flags := mysql.CURSOR_TYPE_NO_CURSOR
239-
if paramsNum > 0 {
243+
if s.conn.capability&mysql.CLIENT_QUERY_ATTRIBUTES > 0 && paramsNum+qaLen > 0 {
240244
flags |= mysql.PARAMETER_COUNT_AVAILABLE
241245
}
242246
data.WriteByte(flags)
@@ -246,7 +250,7 @@ func (s *Stmt) write(args ...any) error {
246250

247251
if paramsNum > 0 || (s.conn.capability&mysql.CLIENT_QUERY_ATTRIBUTES > 0 && (flags&mysql.PARAMETER_COUNT_AVAILABLE > 0)) {
248252
if s.conn.capability&mysql.CLIENT_QUERY_ATTRIBUTES > 0 {
249-
paramsNum += len(s.conn.queryAttributes)
253+
paramsNum += qaLen
250254
data.Write(mysql.PutLengthEncodedInt(uint64(paramsNum)))
251255
}
252256
if paramsNum > 0 {

0 commit comments

Comments
 (0)