@@ -54,13 +54,14 @@ var _ compatibleSource = &mssql.Source{}
5454var compatibleSources = [... ]string {cloudsqlmssql .SourceKind , mssql .SourceKind }
5555
5656type Config struct {
57- Name string `yaml:"name" validate:"required"`
58- Kind string `yaml:"kind" validate:"required"`
59- Source string `yaml:"source" validate:"required"`
60- Description string `yaml:"description" validate:"required"`
61- Statement string `yaml:"statement" validate:"required"`
62- AuthRequired []string `yaml:"authRequired"`
63- Parameters tools.Parameters `yaml:"parameters"`
57+ Name string `yaml:"name" validate:"required"`
58+ Kind string `yaml:"kind" validate:"required"`
59+ Source string `yaml:"source" validate:"required"`
60+ Description string `yaml:"description" validate:"required"`
61+ Statement string `yaml:"statement" validate:"required"`
62+ AuthRequired []string `yaml:"authRequired"`
63+ Parameters tools.Parameters `yaml:"parameters"`
64+ TemplateParameters tools.Parameters `yaml:"templateParameters"`
6465}
6566
6667// validate interface
@@ -83,22 +84,26 @@ func (cfg Config) Initialize(srcs map[string]sources.Source) (tools.Tool, error)
8384 return nil , fmt .Errorf ("invalid source for %q tool: source kind must be one of %q" , kind , compatibleSources )
8485 }
8586
87+ allParameters , paramManifest , paramMcpManifest := tools .ProcessParameters (cfg .TemplateParameters , cfg .Parameters )
88+
8689 mcpManifest := tools.McpManifest {
8790 Name : cfg .Name ,
8891 Description : cfg .Description ,
89- InputSchema : cfg . Parameters . McpManifest () ,
92+ InputSchema : paramMcpManifest ,
9093 }
9194
9295 // finish tool setup
9396 t := Tool {
94- Name : cfg .Name ,
95- Kind : kind ,
96- Parameters : cfg .Parameters ,
97- Statement : cfg .Statement ,
98- AuthRequired : cfg .AuthRequired ,
99- Db : s .MSSQLDB (),
100- manifest : tools.Manifest {Description : cfg .Description , Parameters : cfg .Parameters .Manifest (), AuthRequired : cfg .AuthRequired },
101- mcpManifest : mcpManifest ,
97+ Name : cfg .Name ,
98+ Kind : kind ,
99+ Parameters : cfg .Parameters ,
100+ TemplateParameters : cfg .TemplateParameters ,
101+ AllParams : allParameters ,
102+ Statement : cfg .Statement ,
103+ AuthRequired : cfg .AuthRequired ,
104+ Db : s .MSSQLDB (),
105+ manifest : tools.Manifest {Description : cfg .Description , Parameters : paramManifest , AuthRequired : cfg .AuthRequired },
106+ mcpManifest : mcpManifest ,
102107 }
103108 return t , nil
104109}
@@ -107,10 +112,12 @@ func (cfg Config) Initialize(srcs map[string]sources.Source) (tools.Tool, error)
107112var _ tools.Tool = Tool {}
108113
109114type Tool struct {
110- Name string `yaml:"name"`
111- Kind string `yaml:"kind"`
112- AuthRequired []string `yaml:"authRequired"`
113- Parameters tools.Parameters `yaml:"parameters"`
115+ Name string `yaml:"name"`
116+ Kind string `yaml:"kind"`
117+ AuthRequired []string `yaml:"authRequired"`
118+ Parameters tools.Parameters `yaml:"parameters"`
119+ TemplateParameters tools.Parameters `yaml:"templateParameters"`
120+ AllParams tools.Parameters `yaml:"allParams"`
114121
115122 Db * sql.DB
116123 Statement string
@@ -119,18 +126,29 @@ type Tool struct {
119126}
120127
121128func (t Tool ) Invoke (ctx context.Context , params tools.ParamValues ) ([]any , error ) {
122- namedArgs := make ([]any , 0 , len (params ))
123- paramsMap := params .AsReversedMap ()
129+ paramsMap := params .AsMap ()
130+ newStatement , err := tools .ResolveTemplateParams (t .TemplateParameters , t .Statement , paramsMap )
131+ if err != nil {
132+ return nil , fmt .Errorf ("unable to extract template params %w" , err )
133+ }
134+
135+ newParams , err := tools .GetParams (t .Parameters , paramsMap )
136+ if err != nil {
137+ return nil , fmt .Errorf ("unable to extract standard params %w" , err )
138+ }
139+
140+ namedArgs := make ([]any , 0 , len (newParams ))
141+ newParamsMap := newParams .AsReversedMap ()
124142 // To support both named args (e.g @id) and positional args (e.g @p1), check if arg name is contained in the statement.
125- for _ , v := range params .AsSlice () {
126- paramName := paramsMap [v ]
127- if strings .Contains (t . Statement , "@" + paramName ) {
143+ for _ , v := range newParams .AsSlice () {
144+ paramName := newParamsMap [v ]
145+ if strings .Contains (newStatement , "@" + paramName ) {
128146 namedArgs = append (namedArgs , sql .Named (paramName , v ))
129147 } else {
130148 namedArgs = append (namedArgs , v )
131149 }
132150 }
133- rows , err := t .Db .QueryContext (ctx , t . Statement , namedArgs ... )
151+ rows , err := t .Db .QueryContext (ctx , newStatement , namedArgs ... )
134152 if err != nil {
135153 return nil , fmt .Errorf ("unable to execute query: %w" , err )
136154 }
@@ -173,7 +191,7 @@ func (t Tool) Invoke(ctx context.Context, params tools.ParamValues) ([]any, erro
173191}
174192
175193func (t Tool ) ParseParams (data map [string ]any , claims map [string ]map [string ]any ) (tools.ParamValues , error ) {
176- return tools .ParseParams (t .Parameters , data , claims )
194+ return tools .ParseParams (t .AllParams , data , claims )
177195}
178196
179197func (t Tool ) Manifest () tools.Manifest {
0 commit comments