Skip to content

Commit 6512704

Browse files
feat: Add support for read-only in Spanner tool (#563)
Allowing user to add `readOnly` field in spanner tools. The existing tool doesn't work for reading schema tables since schema tables can only be accessed through read-only transaction. This PR also resolve #435 for Spanner tool. --------- Co-authored-by: Yuan <45984206+Yuan325@users.noreply.github.com>
1 parent 04dcf47 commit 6512704

3 files changed

Lines changed: 79 additions & 28 deletions

File tree

‎docs/en/resources/tools/spanner-sql.md‎

Lines changed: 1 addition & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -124,3 +124,4 @@ tools:
124124
| description | string | true | Description of the tool that is passed to the LLM. |
125125
| statement | string | true | SQL statement to execute on. |
126126
| parameters | [parameters](_index#specifying-parameters) | false | List of [parameters](_index#specifying-parameters) that will be inserted into the SQL statement. |
127+
| readOnly | bool | false | When set to `true`, the `statement` is run as a read-only transaction. Default: `false`. |

‎internal/tools/spanner/spanner.go‎

Lines changed: 47 additions & 28 deletions
Original file line numberDiff line numberDiff line change
@@ -44,6 +44,7 @@ type Config struct {
4444
Source string `yaml:"source" validate:"required"`
4545
Description string `yaml:"description" validate:"required"`
4646
Statement string `yaml:"statement" validate:"required"`
47+
ReadOnly bool `yaml:"readOnly"`
4748
AuthRequired []string `yaml:"authRequired"`
4849
Parameters tools.Parameters `yaml:"parameters"`
4950
}
@@ -81,6 +82,7 @@ func (cfg Config) Initialize(srcs map[string]sources.Source) (tools.Tool, error)
8182
Parameters: cfg.Parameters,
8283
Statement: cfg.Statement,
8384
AuthRequired: cfg.AuthRequired,
85+
ReadOnly: cfg.ReadOnly,
8486
Client: s.SpannerClient(),
8587
dialect: s.DatabaseDialect(),
8688
manifest: tools.Manifest{Description: cfg.Description, Parameters: cfg.Parameters.Manifest(), AuthRequired: cfg.AuthRequired},
@@ -97,7 +99,7 @@ type Tool struct {
9799
Kind string `yaml:"kind"`
98100
AuthRequired []string `yaml:"authRequired"`
99101
Parameters tools.Parameters `yaml:"parameters"`
100-
102+
ReadOnly bool `yaml:"readOnly"`
101103
Client *spanner.Client
102104
dialect string
103105
Statement string
@@ -116,45 +118,62 @@ func getMapParams(params tools.ParamValues, dialect string) (map[string]interfac
116118
}
117119
}
118120

121+
// processRows iterates over the spanner.RowIterator and converts each row to a map[string]any.
122+
func processRows(iter *spanner.RowIterator) ([]any, error) {
123+
var out []any
124+
defer iter.Stop()
125+
126+
for {
127+
row, err := iter.Next()
128+
if err == iterator.Done {
129+
break
130+
}
131+
if err != nil {
132+
return nil, fmt.Errorf("unable to parse row: %w", err)
133+
}
134+
135+
vMap := make(map[string]any)
136+
cols := row.ColumnNames()
137+
for i, c := range cols {
138+
vMap[c] = row.ColumnValue(i)
139+
}
140+
out = append(out, vMap)
141+
}
142+
return out, nil
143+
}
144+
119145
func (t Tool) Invoke(ctx context.Context, params tools.ParamValues) ([]any, error) {
120146
mapParams, err := getMapParams(params, t.dialect)
121147
if err != nil {
122148
return nil, fmt.Errorf("fail to get map params: %w", err)
123149
}
124150

125-
var out []any
126-
127-
_, err = t.Client.ReadWriteTransaction(ctx, func(ctx context.Context, txn *spanner.ReadWriteTransaction) error {
128-
stmt := spanner.Statement{
129-
SQL: t.Statement,
130-
Params: mapParams,
131-
}
132-
iter := txn.Query(ctx, stmt)
133-
defer iter.Stop()
151+
var results []any
152+
var opErr error
153+
stmt := spanner.Statement{
154+
SQL: t.Statement,
155+
Params: mapParams,
156+
}
134157

135-
for {
136-
row, err := iter.Next()
137-
if err == iterator.Done {
138-
return nil
139-
}
158+
if t.ReadOnly {
159+
iter := t.Client.Single().Query(ctx, stmt)
160+
results, opErr = processRows(iter)
161+
} else {
162+
_, opErr = t.Client.ReadWriteTransaction(ctx, func(ctx context.Context, txn *spanner.ReadWriteTransaction) error {
163+
iter := txn.Query(ctx, stmt)
164+
results, err = processRows(iter)
140165
if err != nil {
141-
return fmt.Errorf("unable to parse row: %w", err)
142-
}
143-
144-
vMap := make(map[string]any)
145-
cols := row.ColumnNames()
146-
for i, c := range cols {
147-
vMap[c] = row.ColumnValue(i)
166+
return err
148167
}
168+
return nil
169+
})
170+
}
149171

150-
out = append(out, vMap)
151-
}
152-
})
153-
if err != nil {
154-
return nil, fmt.Errorf("unable to execute client: %w", err)
172+
if opErr != nil {
173+
return nil, fmt.Errorf("unable to execute client: %w", opErr)
155174
}
156175

157-
return out, nil
176+
return results, nil
158177
}
159178

160179
func (t Tool) ParseParams(data map[string]any, claims map[string]map[string]any) (tools.ParamValues, error) {

‎internal/tools/spanner/spanner_test.go‎

Lines changed: 31 additions & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -64,6 +64,37 @@ func TestParseFromYamlSpanner(t *testing.T) {
6464
},
6565
},
6666
},
67+
{
68+
desc: "read only set to true",
69+
in: `
70+
tools:
71+
example_tool:
72+
kind: spanner-sql
73+
source: my-pg-instance
74+
description: some description
75+
readOnly: true
76+
statement: |
77+
SELECT * FROM SQL_STATEMENT;
78+
parameters:
79+
- name: country
80+
type: string
81+
description: some description
82+
`,
83+
want: server.ToolConfigs{
84+
"example_tool": spanner.Config{
85+
Name: "example_tool",
86+
Kind: spanner.ToolKind,
87+
Source: "my-pg-instance",
88+
Description: "some description",
89+
Statement: "SELECT * FROM SQL_STATEMENT;\n",
90+
ReadOnly: true,
91+
AuthRequired: []string{},
92+
Parameters: []tools.Parameter{
93+
tools.NewStringParameter("country", "some description"),
94+
},
95+
},
96+
},
97+
},
6798
}
6899
for _, tc := range tcs {
69100
t.Run(tc.desc, func(t *testing.T) {

0 commit comments

Comments
 (0)