Skip to content
Merged
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
1 change: 0 additions & 1 deletion mcp/cmd/dolt-mcp-server/main.go
Original file line number Diff line number Diff line change
Expand Up @@ -57,7 +57,6 @@ var httpKeyFile = flag.String(httpKeyFlag, "", "Path to TLS private key file for
var httpCAFile = flag.String(httpCAFlag, "", "Path to TLS CA certificate file for HTTPS. If provided, all TLS parameters must be provided otherwise it will be ignored.")
var jwkClaims = flag.String(jwkClaimsFlag, "", "A comma-separated list of key=value pairs for JWT claims for authentication.")
var jwkURL = flag.String(jwkURLFlag, "", "The URL of the JWKS server for JWT authentication.")

var help = flag.Bool(helpFlag, false, "If true, prints Dolt MCP server help information.")
var version = flag.Bool(versionFlag, false, "If true, prints the Dolt MCP server version.")

Expand Down
61 changes: 12 additions & 49 deletions mcp/pkg/db/config.go
Original file line number Diff line number Diff line change
Expand Up @@ -2,8 +2,6 @@ package db

import (
"errors"
"fmt"
"strings"
)

var ErrNoHostDefined = errors.New("no host defined")
Expand All @@ -12,52 +10,17 @@ var ErrNoDatabaseNameDefined = errors.New("no database name defined")
var ErrNoPortDefined = errors.New("no port defined")

type Config struct {
DSN string `yaml:"dsn" json:"dsn"`
Host string `yaml:"host" json:"host"`
User string `yaml:"user" json:"user"`
Password string `yaml:"password" json:"password"`
DatabaseName string `yaml:"database_name" json:"database_name"`
Port int `yaml:"port" json:"port"`
ParseTime bool `yaml:"parse_time" json:"parse_time"`
MultiStatements bool `yaml:"multi_statements" json:"multi_statements"`
TLS string `yaml:"tls" json:"tls"`
TLSCAFile string `yaml:"tls_ca_file" json:"tls_ca_file"`
}

func (c *Config) getDSNOptions() string {
options := []string{}

if c.ParseTime {
options = append(options, "parseTime=true")
}

if c.MultiStatements {
options = append(options, "multiStatements=true")
}

if c.TLS != "" {
options = append(options, fmt.Sprintf("tls=%s", c.TLS))
}

if len(options) > 0 {
return "?" + strings.Join(options, "&")
}

return ""
}

func (c *Config) GetDSN() string {
if c.DSN != "" {
return c.DSN
}

dsn := fmt.Sprintf("%s:%s@tcp(%s:%d)/", c.User, c.Password, c.Host, c.Port)
if c.DatabaseName != "" {
dsn += c.DatabaseName
}

dsn += c.getDSNOptions()
return dsn
DSN string `yaml:"dsn" json:"dsn"`
Host string `yaml:"host" json:"host"`
User string `yaml:"user" json:"user"`
Password string `yaml:"password" json:"password"`
DatabaseName string `yaml:"database_name" json:"database_name"`
Port int `yaml:"port" json:"port"`
ParseTime bool `yaml:"parse_time" json:"parse_time"`
MultiStatements bool `yaml:"multi_statements" json:"multi_statements"`
TLS string `yaml:"tls" json:"tls"`
TLSCAFile string `yaml:"tls_ca_file" json:"tls_ca_file"`
DialectType DialectType `yaml:"dialect_type" json:"dialect_type"`
}

func (c *Config) Validate() error {
Expand All @@ -73,4 +36,4 @@ func (c *Config) Validate() error {
}
}
return nil
}
}
33 changes: 7 additions & 26 deletions mcp/pkg/db/database.go
Original file line number Diff line number Diff line change
Expand Up @@ -2,16 +2,11 @@ package db

import (
"context"
"crypto/tls"
"crypto/x509"
"database/sql"
"encoding/csv"
"errors"
"fmt"
"os"
"strings"

"github.com/go-sql-driver/mysql"
)

type ResultFormat int
Expand Down Expand Up @@ -250,29 +245,15 @@ func (d *databaseTransactionImpl) Commit(ctx context.Context) (err error) {
}

func newDB(config Config) (*sql.DB, error) {
// If a CA file is provided, register a custom TLS config
if config.TLSCAFile != "" {
rootCertPool := x509.NewCertPool()
pem, err := os.ReadFile(config.TLSCAFile)
if err != nil {
return nil, fmt.Errorf("failed to read CA file %s: %w", config.TLSCAFile, err)
}
if ok := rootCertPool.AppendCertsFromPEM(pem); !ok {
return nil, fmt.Errorf("failed to append CA certificate from %s", config.TLSCAFile)
}
tlsConfig := &tls.Config{
RootCAs: rootCertPool,
}
if err := mysql.RegisterTLSConfig("custom", tlsConfig); err != nil {
return nil, fmt.Errorf("failed to register TLS config: %w", err)
}
// Override the TLS setting to use our custom config
config.TLS = "custom"
dialect := NewDialect(config.DialectType)

if err := dialect.ConfigureTLS(&config); err != nil {
return nil, err
}

dsn := config.GetDSN()
dsn := dialect.FormatDSN(config)

db, err := sql.Open("mysql", dsn)
db, err := sql.Open(dialect.DriverName(), dsn)
if err != nil {
return nil, err
}
Expand All @@ -282,4 +263,4 @@ func newDB(config Config) (*sql.DB, error) {
}

return db, nil
}
}
68 changes: 68 additions & 0 deletions mcp/pkg/db/dialect.go
Original file line number Diff line number Diff line change
@@ -0,0 +1,68 @@
package db

import "errors"

// DialectType represents the type of SQL database dialect.
type DialectType string

const (
DialectMySQL DialectType = "mysql"
DialectPostgres DialectType = "postgres"
)

// Validation errors returned by Dialect validation methods.
var (
ErrInvalidSQLReadQuery = errors.New("invalid read query")
ErrInvalidSQLWriteQuery = errors.New("invalid write query")
ErrInvalidCreateTableSQLQuery = errors.New("invalid create table statement")
ErrInvalidAlterTableSQLQuery = errors.New("invalid alter table statement")
)

// DoltProcedure represents a Dolt stored procedure name shared across dialects.
type DoltProcedure string

const (
DoltCheckout DoltProcedure = "DOLT_CHECKOUT"
DoltCommit DoltProcedure = "DOLT_COMMIT"
DoltBranch DoltProcedure = "DOLT_BRANCH"
DoltAdd DoltProcedure = "DOLT_ADD"
DoltReset DoltProcedure = "DOLT_RESET"
DoltMerge DoltProcedure = "DOLT_MERGE"
DoltRemote DoltProcedure = "DOLT_REMOTE"
DoltClone DoltProcedure = "DOLT_CLONE"
DoltFetch DoltProcedure = "DOLT_FETCH"
DoltPush DoltProcedure = "DOLT_PUSH"
DoltPull DoltProcedure = "DOLT_PULL"
)

// Dialect encapsulates all SQL dialect differences between database engines.
type Dialect interface {
// SupportsTool returns whether the given tool name is supported by this dialect.
SupportsTool(toolName string) bool

// Connection setup
DriverName() string
FormatDSN(c Config) string
ConfigureTLS(c *Config) error

// SQL generation
QuoteIdentifier(name string) string
CallProcedure(proc DoltProcedure, args ...string) string
UseDatabase(database string) string

// SQL validation
ValidateReadQuery(query string) error
ValidateWriteQuery(query string) error
ValidateCreateTableQuery(query string) error
ValidateAlterTableQuery(query string) error
}

// NewDialect creates a Dialect for the given DialectType.
func NewDialect(dt DialectType) Dialect {
switch dt {
case DialectPostgres:
panic("postgres dialect not yet implemented")
default:
return &MySQLDialect{}
}
}
164 changes: 164 additions & 0 deletions mcp/pkg/db/dialect_mysql.go
Original file line number Diff line number Diff line change
@@ -0,0 +1,164 @@
package db

import (
"crypto/tls"
"crypto/x509"
"fmt"
"os"
"strings"

"github.com/go-sql-driver/mysql"

gosql "github.com/dolthub/go-mysql-server/sql"
"github.com/dolthub/vitess/go/vt/sqlparser"
)

// MySQLDialect implements Dialect for MySQL-compatible Dolt servers.
type MySQLDialect struct {
unsupportedTools map[string]bool
}

var _ Dialect = &MySQLDialect{}

func (d *MySQLDialect) SupportsTool(toolName string) bool {
if d.unsupportedTools == nil {
return true
}
return !d.unsupportedTools[toolName]
}

func (d *MySQLDialect) DriverName() string {
return "mysql"
}

func (d *MySQLDialect) FormatDSN(c Config) string {
if c.DSN != "" {
return c.DSN
}

dsn := fmt.Sprintf("%s:%s@tcp(%s:%d)/", c.User, c.Password, c.Host, c.Port)
if c.DatabaseName != "" {
dsn += c.DatabaseName
}

options := []string{}
if c.ParseTime {
options = append(options, "parseTime=true")
}
if c.MultiStatements {
options = append(options, "multiStatements=true")
}
if c.TLS != "" {
options = append(options, fmt.Sprintf("tls=%s", c.TLS))
}
if len(options) > 0 {
dsn += "?" + strings.Join(options, "&")
}

return dsn
}

func (d *MySQLDialect) ConfigureTLS(c *Config) error {
if c.TLSCAFile == "" {
return nil
}
rootCertPool := x509.NewCertPool()
pem, err := os.ReadFile(c.TLSCAFile)
if err != nil {
return fmt.Errorf("failed to read CA file %s: %w", c.TLSCAFile, err)
}
if ok := rootCertPool.AppendCertsFromPEM(pem); !ok {
return fmt.Errorf("failed to append CA certificate from %s", c.TLSCAFile)
}
tlsConfig := &tls.Config{
RootCAs: rootCertPool,
}
if err := mysql.RegisterTLSConfig("custom", tlsConfig); err != nil {
return fmt.Errorf("failed to register TLS config: %w", err)
}
c.TLS = "custom"
return nil
}

func (d *MySQLDialect) QuoteIdentifier(name string) string {
return fmt.Sprintf("`%s`", name)
}

func (d *MySQLDialect) CallProcedure(proc DoltProcedure, args ...string) string {
quotedArgs := make([]string, len(args))
for i, arg := range args {
quotedArgs[i] = fmt.Sprintf("'%s'", arg)
}
return fmt.Sprintf("CALL %s(%s);", string(proc), strings.Join(quotedArgs, ", "))
}

func (d *MySQLDialect) UseDatabase(database string) string {
return fmt.Sprintf("USE `%s`;", database)
}

// SQL validation using the Vitess MySQL parser.

func (d *MySQLDialect) parseSQLQuery(query string) (sqlparser.Statement, error) {
sqlCtx := gosql.NewEmptyContext()
sqlMode := gosql.LoadSqlMode(sqlCtx)
return sqlparser.ParseWithOptions(sqlCtx, query, sqlMode.ParserOptions())
}

func (d *MySQLDialect) isReadOnlyStatement(stmt sqlparser.Statement) bool {
switch stmt.(type) {
case sqlparser.SelectStatement:
return true
case *sqlparser.Show:
return true
case *sqlparser.Explain, *sqlparser.OtherRead:
return true
default:
return false
}
}

func (d *MySQLDialect) ValidateReadQuery(query string) error {
stmt, err := d.parseSQLQuery(query)
if err != nil {
return err
}
if d.isReadOnlyStatement(stmt) {
return nil
}
return ErrInvalidSQLReadQuery
}

func (d *MySQLDialect) ValidateWriteQuery(query string) error {
stmt, err := d.parseSQLQuery(query)
if err != nil {
return err
}
if d.isReadOnlyStatement(stmt) {
return ErrInvalidSQLWriteQuery
}
return nil
}

func (d *MySQLDialect) ValidateCreateTableQuery(query string) error {
stmt, err := d.parseSQLQuery(query)
if err != nil {
return err
}
switch stmt.(type) {
case *sqlparser.DDL:
return nil
}
return ErrInvalidCreateTableSQLQuery
}

func (d *MySQLDialect) ValidateAlterTableQuery(query string) error {
stmt, err := d.parseSQLQuery(query)
if err != nil {
return err
}
switch stmt.(type) {
case *sqlparser.AlterTable:
return nil
}
return ErrInvalidAlterTableSQLQuery
}
Loading
Loading