diff --git a/go.mod b/go.mod index 8cc60d4..614d15e 100644 --- a/go.mod +++ b/go.mod @@ -3,7 +3,6 @@ module github.com/shadow1ng/fscan go 1.20 require ( - github.com/denisenkom/go-mssqldb v0.12.3 github.com/fatih/color v1.18.0 github.com/go-ldap/ldap/v3 v3.4.9 github.com/go-sql-driver/mysql v1.8.1 @@ -41,8 +40,6 @@ require ( github.com/antlr/antlr4/runtime/Go/antlr v1.4.10 // indirect github.com/geoffgarside/ber v1.1.0 // indirect github.com/go-asn1-ber/asn1-ber v1.5.7 // indirect - github.com/golang-sql/civil v0.0.0-20190719163853-cb61b32ac6fe // indirect - github.com/golang-sql/sqlexp v0.1.0 // indirect github.com/hashicorp/errwrap v1.0.0 // indirect github.com/hashicorp/go-multierror v1.1.1 // indirect github.com/hashicorp/go-uuid v1.0.3 // indirect diff --git a/go.sum b/go.sum index 4b7216e..acbe998 100644 --- a/go.sum +++ b/go.sum @@ -1,9 +1,6 @@ cloud.google.com/go v0.26.0/go.mod h1:aQUYkXzVsufM+DwF1aE+0xfcU+56JwCaLick0ClmMTw= filippo.io/edwards25519 v1.1.0 h1:FNf4tywRC1HmFuKW5xopWpigGjJKiJSV0Cqo0cJWDaA= filippo.io/edwards25519 v1.1.0/go.mod h1:BxyFTGdWcka3PhytdK4V28tE5sGfRvvvRV7EaN4VDT4= -github.com/Azure/azure-sdk-for-go/sdk/azcore v0.19.0/go.mod h1:h6H6c8enJmmocHUbLiiGY6sx7f9i+X3m1CHdd5c6Rdw= -github.com/Azure/azure-sdk-for-go/sdk/azidentity v0.11.0/go.mod h1:HcM1YX14R7CJcghJGOYCgdezslRSVzqwLf/q+4Y2r/0= -github.com/Azure/azure-sdk-for-go/sdk/internal v0.7.0/go.mod h1:yqy467j36fJxcRV2TzfVZ1pCb5vxm4BtZPUdYWe/Xo8= github.com/Azure/go-ntlmssp v0.0.0-20221128193559-754e69321358 h1:mFRzDkZVAjdal+s7s0MwaRv9igoPqLRdzOLzw/8Xvq8= github.com/Azure/go-ntlmssp v0.0.0-20221128193559-754e69321358/go.mod h1:chxPXzSsl7ZWRAuOIE23GDNzjWuZquvFlgA8xmpunjU= github.com/BurntSushi/toml v0.3.1/go.mod h1:xHWCNGjB5oqiDr8zfno3MHue2Ht5sIBksp03qcyfWMU= @@ -19,9 +16,6 @@ github.com/creack/pty v1.1.9/go.mod h1:oKZEueFk5CKHvIhNR5MUki03XCEU+Q6VDXinZuGJ3 github.com/davecgh/go-spew v1.1.0/go.mod h1:J7Y8YcW2NihsgmVo/mv3lAwl/skON4iLHjSsI+c5H38= github.com/davecgh/go-spew v1.1.1 h1:vj9j/u1bqnvCEfJOwUhtlOARqs3+rkHYY13jYWTU97c= github.com/davecgh/go-spew v1.1.1/go.mod h1:J7Y8YcW2NihsgmVo/mv3lAwl/skON4iLHjSsI+c5H38= -github.com/denisenkom/go-mssqldb v0.12.3 h1:pBSGx9Tq67pBOTLmxNuirNTeB8Vjmf886Kx+8Y+8shw= -github.com/denisenkom/go-mssqldb v0.12.3/go.mod h1:k0mtMFOnU+AihqFxPMiF05rtiDrorD1Vrm1KEz5hxDo= -github.com/dnaeon/go-vcr v1.2.0/go.mod h1:R4UdLID7HZT3taECzJs4YgbbH6PIGXB6W/sc5OLb6RQ= github.com/envoyproxy/go-control-plane v0.9.0/go.mod h1:YTl/9mNaCwkRvm6d1a2C3ymFceY/DCBVvsKhRF0iEA4= github.com/envoyproxy/go-control-plane v0.9.4/go.mod h1:6rpuAdCZL397s3pYoYcLgu1mIlRU8Am5FuJP05cCM98= github.com/envoyproxy/protoc-gen-validate v0.1.0/go.mod h1:iSmxcyjqTsJpI2R4NaDN7+kN2VEUnK/pcBlmesArF7c= @@ -35,10 +29,6 @@ github.com/go-ldap/ldap/v3 v3.4.9 h1:KxX9eO44/MpqPXVVMPJDB+k/35GEePHE/Jfvl7oRMUo github.com/go-ldap/ldap/v3 v3.4.9/go.mod h1:+CE/4PPOOdEPGTi2B7qXKQOq+pNBvXZtlBNcVZY0AWI= github.com/go-sql-driver/mysql v1.8.1 h1:LedoTUt/eveggdHS9qUFC1EFSa8bU2+1pZjSRpvNJ1Y= github.com/go-sql-driver/mysql v1.8.1/go.mod h1:wEBSXgmK//2ZFJyE+qWnIsVGmvmEKlqwuVSjsCm7DZg= -github.com/golang-sql/civil v0.0.0-20190719163853-cb61b32ac6fe h1:lXe2qZdvpiX5WZkZR4hgp4KJVfY3nMkvmwbVkpv1rVY= -github.com/golang-sql/civil v0.0.0-20190719163853-cb61b32ac6fe/go.mod h1:8vg3r2VgvsThLBIFL93Qb5yWzgyZWhEmBwUJWevAkK0= -github.com/golang-sql/sqlexp v0.1.0 h1:ZCD6MBpcuOVfGVqsEmY5/4FtYiKz6tSyUv9LPEDei6A= -github.com/golang-sql/sqlexp v0.1.0/go.mod h1:J4ad9Vo8ZCWQ2GMrC4UCQy1JpCbwU9m3EOqtpKwwwHI= github.com/golang/glog v0.0.0-20160126235308-23def4e6c14b/go.mod h1:SBH7ygxi8pfUlaOkMMuAQtPIUF8ecWP5IEl/CR7VP2Q= github.com/golang/mock v1.1.1/go.mod h1:oTYuIxOrZwtPieC+H1uAHpcLFnEyAGVDL/k47Jfbm0A= github.com/golang/protobuf v1.2.0/go.mod h1:6lQm79b+lXiMfvg/cZm0SGofjICqVBUtrP5yJMmIC1U= @@ -118,12 +108,10 @@ github.com/mattn/go-isatty v0.0.20 h1:xfD0iDuEKnDkl03q4limB+vH+GxLEtL/jb4xVJSWWE github.com/mattn/go-isatty v0.0.20/go.mod h1:W+V8PltTTMOvKvAeJH7IuucS94S2C6jfK/D7dTCTo3Y= github.com/mitchellh/go-vnc v0.0.0-20150629162542-723ed9867aed h1:FI2NIv6fpef6BQl2u3IZX/Cj20tfypRF4yd+uaHOMtI= github.com/mitchellh/go-vnc v0.0.0-20150629162542-723ed9867aed/go.mod h1:3rdaFaCv4AyBgu5ALFM0+tSuHrBh6v692nyQe3ikrq0= -github.com/modocache/gover v0.0.0-20171022184752-b58185e213c5/go.mod h1:caMODM3PzxT8aQXRPkAt8xlV/e7d7w8GM5g0fa5F0D8= github.com/nicksnyder/go-i18n/v2 v2.4.0 h1:3IcvPOAvnCKwNm0TB0dLDTuawWEj+ax/RERNC+diLMM= github.com/nicksnyder/go-i18n/v2 v2.4.0/go.mod h1:nxYSZE9M0bf3Y70gPQjN9ha7XNHX7gMc814+6wVyEI4= github.com/panjf2000/ants/v2 v2.11.3 h1:AfI0ngBoXJmYOpDh9m516vjqoUu2sLrIVgppI9TZVpg= github.com/panjf2000/ants/v2 v2.11.3/go.mod h1:8u92CYMUc6gyvTIw8Ru7Mt7+/ESnJahz5EVtqfrilek= -github.com/pkg/browser v0.0.0-20180916011732-0a3d74bf9ce4/go.mod h1:4OwLy04Bl9Ef3GJJCoec+30X3LQs/0/m4HFRt/2LUSA= github.com/pkg/errors v0.9.1 h1:FEBLx1zS214owpjy7qsBeixbURkuhQAwrK5UwLGTwt4= github.com/pkg/errors v0.9.1/go.mod h1:bwawxfHBFNV+L2hUp1rHADufV3IMtnDRdf1r5NINEl0= github.com/pmezard/go-difflib v1.0.0 h1:4DBwDE0NGyQoBHbLQYPwSUPoCMWR5BEzIk/f1lZbAQM= @@ -141,7 +129,6 @@ github.com/stretchr/objx v0.4.0/go.mod h1:YvHI0jy2hoMjB+UWwv71VJQ9isScKT/TqJzVSS github.com/stretchr/objx v0.5.0/go.mod h1:Yh+to48EsGEfYuaHDzXPcE3xhTkx73EhmCGUpEOglKo= github.com/stretchr/testify v1.4.0/go.mod h1:j7eGeouHqKxXV5pUuKE4zz7dFj8WfuZ+81PSLYec5m4= github.com/stretchr/testify v1.5.1/go.mod h1:5W2xD1RspED5o8YsWQXVCued0rvSQ+mT+I5cxcmMvtA= -github.com/stretchr/testify v1.7.0/go.mod h1:6Fq8oRcR53rry900zMqJjRRixrwX3KX962/h/Wwjteg= github.com/stretchr/testify v1.7.1/go.mod h1:6Fq8oRcR53rry900zMqJjRRixrwX3KX962/h/Wwjteg= github.com/stretchr/testify v1.8.0/go.mod h1:yNjHg4UonilssWZ8iaSj1OCr/vHnekPRkoO+kdMU+MU= github.com/stretchr/testify v1.8.1/go.mod h1:w2LPCIKwWwSfY2zedu0+kehJoqGctiVI29o6fzry7u4= @@ -155,9 +142,7 @@ golang.org/x/crypto v0.0.0-20190308221718-c2843e01d9a2/go.mod h1:djNgcEr1/C05ACk golang.org/x/crypto v0.0.0-20200622213623-75b288015ac9/go.mod h1:LzIPMQfyMNhhGPhUkYOs5KpL4U8rLKemX1yGLhDgUto= golang.org/x/crypto v0.0.0-20200728195943-123391ffb6de/go.mod h1:LzIPMQfyMNhhGPhUkYOs5KpL4U8rLKemX1yGLhDgUto= golang.org/x/crypto v0.0.0-20201012173705-84dcc777aaee/go.mod h1:LzIPMQfyMNhhGPhUkYOs5KpL4U8rLKemX1yGLhDgUto= -golang.org/x/crypto v0.0.0-20201016220609-9e8e0b390897/go.mod h1:LzIPMQfyMNhhGPhUkYOs5KpL4U8rLKemX1yGLhDgUto= golang.org/x/crypto v0.0.0-20210921155107-089bfa567519/go.mod h1:GvvjBRRGRdwPK5ydBHafDWAxML/pGHZbMvKqRZ5+Abc= -golang.org/x/crypto v0.0.0-20220622213112-05595931fe9d/go.mod h1:IxCIyHEi3zRg3s0A5j5BB6A9Jmi73HwBIUl50j+osU4= golang.org/x/crypto v0.6.0/go.mod h1:OFC/31mSvZgRz0V1QTNCzfAI1aIRzbiufJtkMIlEp58= golang.org/x/crypto v0.13.0/go.mod h1:y6Z2r+Rw4iayiXXAIxJIDAJ1zMW4yaTpebo8fPOliYc= golang.org/x/crypto v0.19.0/go.mod h1:Iy9bg/ha4yyC70EfRS8jz+B6ybOBKMaSxLj6P6oBDfU= @@ -183,8 +168,6 @@ golang.org/x/net v0.0.0-20190620200207-3b0461eec859/go.mod h1:z5CRVTTTmAJ677TzLL golang.org/x/net v0.0.0-20200114155413-6afb5195e5aa/go.mod h1:z5CRVTTTmAJ677TzLLGU+0bjPO0LkuOLi4/5GtJWs/s= golang.org/x/net v0.0.0-20201010224723-4f7140c49acb/go.mod h1:sp8m0HH+o8qH0wwXwYZr8TS3Oi6o0r6Gce1SSxlDquU= golang.org/x/net v0.0.0-20210226172049-e18ecbb05110/go.mod h1:m0MpNAwzfU5UDzcl9v0D8zg8gWTRqZa9RBIspLL5mdg= -golang.org/x/net v0.0.0-20210610132358-84b48f89b13b/go.mod h1:9nx3DQGgdP8bBQD5qxJ1jj9UTztislL4KSBs9R2vV5Y= -golang.org/x/net v0.0.0-20211112202133-69e39bad7dc2/go.mod h1:9nx3DQGgdP8bBQD5qxJ1jj9UTztislL4KSBs9R2vV5Y= golang.org/x/net v0.0.0-20220722155237-a158d28d115b/go.mod h1:XRhObCWvk6IyKnWLug+ECip1KBveYUHfp+8e9klMJ9c= golang.org/x/net v0.6.0/go.mod h1:2Tu9+aMcznHK/AK1HMvgo6xiTLG5rD5rZLDS+rp2Bjs= golang.org/x/net v0.7.0/go.mod h1:2Tu9+aMcznHK/AK1HMvgo6xiTLG5rD5rZLDS+rp2Bjs= @@ -211,7 +194,6 @@ golang.org/x/sys v0.0.0-20190215142949-d0b11bdaac8a/go.mod h1:STP8DvDyc/dI5b8T5h golang.org/x/sys v0.0.0-20190412213103-97732733099d/go.mod h1:h1NjWce9XRLGQEsW7wpKNCjG9DtNlClVuFLEZdDNbEs= golang.org/x/sys v0.0.0-20200930185726-fdedc70b468f/go.mod h1:h1NjWce9XRLGQEsW7wpKNCjG9DtNlClVuFLEZdDNbEs= golang.org/x/sys v0.0.0-20201119102817-f84b799fce68/go.mod h1:h1NjWce9XRLGQEsW7wpKNCjG9DtNlClVuFLEZdDNbEs= -golang.org/x/sys v0.0.0-20210423082822-04245dca01da/go.mod h1:h1NjWce9XRLGQEsW7wpKNCjG9DtNlClVuFLEZdDNbEs= golang.org/x/sys v0.0.0-20210615035016-665e8c7367d1/go.mod h1:oPkhp1MJrh7nUepCBck5+mAzfO9JrbApNNgaTdGDITg= golang.org/x/sys v0.0.0-20220520151302-bc2c85ada10a/go.mod h1:oPkhp1MJrh7nUepCBck5+mAzfO9JrbApNNgaTdGDITg= golang.org/x/sys v0.0.0-20220722155257-8c9f86f7a55f/go.mod h1:oPkhp1MJrh7nUepCBck5+mAzfO9JrbApNNgaTdGDITg= @@ -236,7 +218,6 @@ golang.org/x/term v0.27.0 h1:WP60Sv1nlK1T6SupCHbXzSaN0b9wUmsPoRS9b61A23Q= golang.org/x/term v0.27.0/go.mod h1:iMsnZpn0cago0GOrHO2+Y7u7JPn5AylBrcoWkElMTSM= golang.org/x/text v0.3.0/go.mod h1:NqM8EUOU14njkJ3fqMW+pc6Ldnwhi/IjpwHt7yyuwOQ= golang.org/x/text v0.3.3/go.mod h1:5Zoc/QRtKVWzQhOtBMvqHzDpF6irO9z98xDceosuGiQ= -golang.org/x/text v0.3.6/go.mod h1:5Zoc/QRtKVWzQhOtBMvqHzDpF6irO9z98xDceosuGiQ= golang.org/x/text v0.3.7/go.mod h1:u+2+/6zg+i71rQMx5EYifcz6MCKuco9NR6JIITiCfzQ= golang.org/x/text v0.7.0/go.mod h1:mrYo+phRRbMaCq/xk9113O4dZlRixOauAjOtrjsXDZ8= golang.org/x/text v0.9.0/go.mod h1:e1OnstbJyHTd6l/uOt8jFFHp6TRDWZR/bV3emEE/zU8= @@ -282,11 +263,9 @@ gopkg.in/check.v1 v1.0.0-20201130134442-10cb98267c6c h1:Hei/4ADfdWqJk1ZMxUNpqntN gopkg.in/check.v1 v1.0.0-20201130134442-10cb98267c6c/go.mod h1:JHkPIbrfpd72SG/EVd6muEfDQjcINNoR0C8j2r3qZ4Q= gopkg.in/errgo.v2 v2.1.0/go.mod h1:hNsd1EY+bozCKY1Ytp96fpM3vjJbqLJn88ws8XvfDNI= gopkg.in/yaml.v2 v2.2.2/go.mod h1:hI93XBmqTisBFMUTm0b8Fm+jr3Dg1NNxqwp+5A1VGuI= -gopkg.in/yaml.v2 v2.2.8/go.mod h1:hI93XBmqTisBFMUTm0b8Fm+jr3Dg1NNxqwp+5A1VGuI= gopkg.in/yaml.v2 v2.4.0 h1:D8xgwECY7CYvx+Y2n4sBz93Jn9JRvxdiyyo8CTfuKaY= gopkg.in/yaml.v2 v2.4.0/go.mod h1:RDklbk79AGWmwhnvt/jBztapEOGDOx6ZbXqjP6csGnQ= gopkg.in/yaml.v3 v3.0.0-20200313102051-9f266ea9e77c/go.mod h1:K4uyk7z7BCEPqu6E+C64Yfv1cQ7kz7rIZviUmN+EgEM= -gopkg.in/yaml.v3 v3.0.0-20210107192922-496545a6307b/go.mod h1:K4uyk7z7BCEPqu6E+C64Yfv1cQ7kz7rIZviUmN+EgEM= gopkg.in/yaml.v3 v3.0.1 h1:fxVm/GzAzEWqLHuvctI91KS9hhNmmWOoWu0XTYJS7CA= gopkg.in/yaml.v3 v3.0.1/go.mod h1:K4uyk7z7BCEPqu6E+C64Yfv1cQ7kz7rIZviUmN+EgEM= honnef.co/go/tools v0.0.0-20190102054323-c2f93a96b099/go.mod h1:rf3lG4BRIbNafJWhAfAdb/ePZxsR/4RtNHQocxwk9r4= diff --git a/plugins/services/mssql.go b/plugins/services/mssql.go index 44490fb..69474cf 100644 --- a/plugins/services/mssql.go +++ b/plugins/services/mssql.go @@ -4,11 +4,9 @@ package services import ( "context" - "database/sql" "fmt" "strings" - _ "github.com/denisenkom/go-mssqldb" // MSSQL driver "github.com/shadow1ng/fscan/common" "github.com/shadow1ng/fscan/common/i18n" "github.com/shadow1ng/fscan/plugins" @@ -65,29 +63,11 @@ func (p *MSSQLPlugin) createAuthFunc(info *common.HostInfo, config *common.Confi // doMSSQLAuth 执行MSSQL认证 func (p *MSSQLPlugin) doMSSQLAuth(ctx context.Context, info *common.HostInfo, cred Credential, config *common.Config, state *common.State) *AuthResult { - connStr := fmt.Sprintf("server=%s;user id=%s;password=%s;port=%d;database=master;encrypt=disable;connection timeout=%d", - info.Host, cred.Username, cred.Password, info.Port, int64(config.Timeout.Seconds())) - - db, err := sql.Open("mssql", connStr) - if err != nil { - state.IncrementTCPFailedPacketCount() - return &AuthResult{ - Success: false, - ErrorType: classifyMSSQLErrorType(err), - Error: err, - } - } - - db.SetConnMaxLifetime(config.Timeout) - db.SetMaxOpenConns(1) - db.SetMaxIdleConns(0) - - pingCtx, cancel := context.WithTimeout(ctx, config.Timeout) + authCtx, cancel := context.WithTimeout(ctx, config.Timeout) defer cancel() - err = db.PingContext(pingCtx) + _, err := mssqlRawLogin(authCtx, info.Host, info.Port, cred.Username, cred.Password, config.Timeout) if err != nil { - _ = db.Close() state.IncrementTCPFailedPacketCount() return &AuthResult{ Success: false, @@ -100,7 +80,6 @@ func (p *MSSQLPlugin) doMSSQLAuth(ctx context.Context, info *common.HostInfo, cr return &AuthResult{ Success: true, - Conn: &SQLDBWrapper{db}, ErrorType: ErrorTypeUnknown, Error: nil, } @@ -148,23 +127,10 @@ func classifyMSSQLErrorType(err error) ErrorType { func (p *MSSQLPlugin) identifyService(ctx context.Context, info *common.HostInfo, config *common.Config, state *common.State) *ScanResult { target := info.Target() - connStr := fmt.Sprintf("server=%s;user id=invalid;password=invalid;port=%d;database=master;encrypt=disable;connection timeout=%d", - info.Host, info.Port, int64(config.Timeout.Seconds())) - - db, err := sql.Open("mssql", connStr) - if err != nil { - return &ScanResult{ - Success: false, - Service: "mssql", - Error: err, - } - } - defer func() { _ = db.Close() }() - - pingCtx, cancel := context.WithTimeout(ctx, config.Timeout) + identifyCtx, cancel := context.WithTimeout(ctx, config.Timeout) defer cancel() - err = db.PingContext(pingCtx) + result, err := mssqlRawLogin(identifyCtx, info.Host, info.Port, "invalid", "invalid", config.Timeout) if err != nil { state.IncrementTCPFailedPacketCount() @@ -178,11 +144,10 @@ func (p *MSSQLPlugin) identifyService(ctx context.Context, info *common.HostInfo errLower = strings.ToLower(err.Error()) } - if err != nil && (strings.Contains(errLower, "login failed") || - strings.Contains(errLower, "mssql") || - strings.Contains(errLower, "sql server")) { - banner = "MSSQL" - } else if err == nil { + if err == nil || (result != nil && result.isMSSQL()) || + (strings.Contains(errLower, "login failed") || + strings.Contains(errLower, "mssql") || + strings.Contains(errLower, "sql server")) { banner = "MSSQL" } else { return &ScanResult{ diff --git a/plugins/services/mssql_raw.go b/plugins/services/mssql_raw.go new file mode 100644 index 0000000..9c2e869 --- /dev/null +++ b/plugins/services/mssql_raw.go @@ -0,0 +1,478 @@ +//go:build plugin_mssql || !plugin_selective + +package services + +import ( + "bytes" + "context" + "encoding/binary" + "fmt" + "io" + "net" + "os" + "sort" + "time" + "unicode/utf16" +) + +const ( + tdsPacketReply = 4 + tdsPacketLogin7 = 16 + tdsPacketPrelogin = 18 + + tdsStatusEOM = 1 + + tdsVersion74 = 0x74000004 + tdsDefaultPacketLen = 4096 + + tdsPreloginVersion = 0 + tdsPreloginEncryption = 1 + tdsPreloginInstOpt = 2 + tdsPreloginThreadID = 3 + tdsPreloginMARS = 4 + tdsPreloginTerminator = 0xff + + tdsEncryptNotSupported = 2 + + tdsTokenError = 0xaa + tdsTokenInfo = 0xab + tdsTokenLoginAck = 0xad + tdsTokenEnvChange = 0xe3 + tdsTokenDone = 0xfd + tdsTokenDoneProc = 0xfe + tdsTokenDoneInProc = 0xff + + tdsDoneError = 0x0002 + tdsDoneSrvError = 0x0100 + + tdsOptionUseDB = 0x20 + tdsOptionSetLang = 0x80 + tdsOptionODBC = 0x02 + + tdsLoginHeaderLen = 94 +) + +type mssqlRawResult struct { + sawPrelogin bool + sawLoginAck bool + errors []mssqlRawError +} + +func (r *mssqlRawResult) isMSSQL() bool { + return r != nil && (r.sawPrelogin || r.sawLoginAck || len(r.errors) > 0) +} + +type mssqlRawError struct { + number int32 + message string +} + +func (e mssqlRawError) Error() string { + if e.message == "" { + return fmt.Sprintf("mssql: error %d", e.number) + } + return "mssql: " + e.message +} + +func mssqlRawLogin(ctx context.Context, host string, port int, username, password string, timeout time.Duration) (*mssqlRawResult, error) { + target := net.JoinHostPort(host, fmt.Sprint(port)) + dialer := net.Dialer{Timeout: timeout} + conn, err := dialer.DialContext(ctx, "tcp", target) + if err != nil { + return nil, err + } + defer conn.Close() + + if deadline, ok := ctx.Deadline(); ok { + _ = conn.SetDeadline(deadline) + } else if timeout > 0 { + _ = conn.SetDeadline(time.Now().Add(timeout)) + } + + result := &mssqlRawResult{} + if err := mssqlSendPrelogin(conn); err != nil { + return result, err + } + if err := mssqlReadPrelogin(conn); err != nil { + return result, err + } + result.sawPrelogin = true + + if err := mssqlSendLogin7(conn, host, username, password); err != nil { + return result, err + } + if err := mssqlReadLoginResponse(conn, result); err != nil { + return result, err + } + if len(result.errors) > 0 { + return result, result.errors[len(result.errors)-1] + } + if !result.sawLoginAck { + return result, fmt.Errorf("mssql: login acknowledgement not received") + } + return result, nil +} + +func mssqlSendPrelogin(w io.Writer) error { + fields := map[byte][]byte{ + tdsPreloginVersion: {0, 0, 0, 0, 0, 0}, + tdsPreloginEncryption: {tdsEncryptNotSupported}, + tdsPreloginInstOpt: {0}, + tdsPreloginThreadID: {0, 0, 0, 0}, + tdsPreloginMARS: {0}, + } + + keys := make([]int, 0, len(fields)) + for k := range fields { + keys = append(keys, int(k)) + } + sort.Ints(keys) + + payload := bytes.NewBuffer(nil) + offset := uint16(len(fields)*5 + 1) + for _, key := range keys { + value := fields[byte(key)] + payload.WriteByte(byte(key)) + _ = binary.Write(payload, binary.BigEndian, offset) + _ = binary.Write(payload, binary.BigEndian, uint16(len(value))) + offset += uint16(len(value)) + } + payload.WriteByte(tdsPreloginTerminator) + for _, key := range keys { + payload.Write(fields[byte(key)]) + } + + return mssqlWritePacket(w, tdsPacketPrelogin, payload.Bytes()) +} + +func mssqlReadPrelogin(r io.Reader) error { + packetType, payload, err := mssqlReadMessage(r) + if err != nil { + return err + } + if packetType != tdsPacketReply { + return fmt.Errorf("mssql: invalid prelogin response packet type %d", packetType) + } + if len(payload) == 0 { + return fmt.Errorf("mssql: empty prelogin response") + } + + fields, err := mssqlParsePreloginFields(payload) + if err != nil { + return err + } + if _, ok := fields[tdsPreloginEncryption]; !ok { + return fmt.Errorf("mssql: prelogin response missing encryption field") + } + return nil +} + +func mssqlParsePreloginFields(payload []byte) (map[byte][]byte, error) { + fields := make(map[byte][]byte) + for pos := 0; ; pos += 5 { + if pos >= len(payload) { + return nil, fmt.Errorf("mssql: invalid prelogin option table") + } + token := payload[pos] + if token == tdsPreloginTerminator { + return fields, nil + } + if pos+5 > len(payload) { + return nil, fmt.Errorf("mssql: truncated prelogin option") + } + offset := int(binary.BigEndian.Uint16(payload[pos+1 : pos+3])) + length := int(binary.BigEndian.Uint16(payload[pos+3 : pos+5])) + if offset < 0 || length < 0 || offset+length > len(payload) { + return nil, fmt.Errorf("mssql: invalid prelogin option bounds") + } + fields[token] = payload[offset : offset+length] + } +} + +func mssqlSendLogin7(w io.Writer, host, username, password string) error { + hostname, _ := os.Hostname() + values := []struct { + text string + password bool + }{ + {hostname, false}, + {username, false}, + {password, true}, + {"fscan", false}, + {host, false}, + {"fscan", false}, + {"", false}, + {"master", false}, + {"", false}, + {"", false}, + } + + encoded := make([][]byte, len(values)) + lengths := make([]uint16, len(values)) + for i, value := range values { + if value.password { + encoded[i] = mssqlEncodePassword(value.text) + } else { + encoded[i] = mssqlUCS2(value.text) + } + lengths[i] = uint16(len(encoded[i]) / 2) + } + + offsets := make([]uint16, len(values)) + offset := uint16(tdsLoginHeaderLen) + for i, value := range encoded { + offsets[i] = offset + offset += uint16(len(value)) + } + + body := bytes.NewBuffer(make([]byte, 0, int(offset))) + put32 := func(v uint32) { _ = binary.Write(body, binary.LittleEndian, v) } + put16 := func(v uint16) { _ = binary.Write(body, binary.LittleEndian, v) } + + put32(uint32(offset)) + put32(tdsVersion74) + put32(tdsDefaultPacketLen) + put32(0) + put32(uint32(os.Getpid())) + put32(0) + body.WriteByte(tdsOptionUseDB | tdsOptionSetLang) + body.WriteByte(tdsOptionODBC) + body.WriteByte(0) + body.WriteByte(0) + put32(0) + put32(0) + + for i := 0; i < 5; i++ { + put16(offsets[i]) + put16(lengths[i]) + } + put16(0) + put16(0) + for i := 5; i < 8; i++ { + put16(offsets[i]) + put16(lengths[i]) + } + body.Write([]byte{0, 0, 0, 0, 0, 0}) + put16(offsets[8]) + put16(0) + for i := 8; i < 10; i++ { + put16(offsets[i]) + put16(lengths[i]) + } + put32(0) + + for _, value := range encoded { + body.Write(value) + } + return mssqlWritePacket(w, tdsPacketLogin7, body.Bytes()) +} + +func mssqlReadLoginResponse(r io.Reader, result *mssqlRawResult) error { + for { + packetType, payload, err := mssqlReadMessage(r) + if err != nil { + return err + } + if packetType != tdsPacketReply { + return fmt.Errorf("mssql: unexpected login response packet type %d", packetType) + } + done, err := mssqlParseLoginTokens(payload, result) + if err != nil { + return err + } + if done || result.sawLoginAck || len(result.errors) > 0 { + return nil + } + } +} + +func mssqlParseLoginTokens(payload []byte, result *mssqlRawResult) (bool, error) { + pos := 0 + for pos < len(payload) { + token := payload[pos] + pos++ + switch token { + case tdsTokenError: + errMsg, next, err := mssqlParseErrorToken(payload, pos) + if err != nil { + return false, err + } + result.errors = append(result.errors, errMsg) + pos = next + case tdsTokenInfo: + next, err := mssqlSkipUSVarError(payload, pos) + if err != nil { + return false, err + } + pos = next + case tdsTokenEnvChange: + next, err := mssqlSkipLen16(payload, pos) + if err != nil { + return false, err + } + pos = next + case tdsTokenLoginAck: + next, err := mssqlSkipLen16(payload, pos) + if err != nil { + return false, err + } + result.sawLoginAck = true + pos = next + case tdsTokenDone, tdsTokenDoneProc, tdsTokenDoneInProc: + if pos+12 > len(payload) { + return false, fmt.Errorf("mssql: truncated done token") + } + status := binary.LittleEndian.Uint16(payload[pos : pos+2]) + pos += 12 + return status&(tdsDoneError|tdsDoneSrvError) == 0, nil + default: + return false, fmt.Errorf("mssql: unexpected login token 0x%02x", token) + } + } + return false, nil +} + +func mssqlParseErrorToken(payload []byte, pos int) (mssqlRawError, int, error) { + if pos+2 > len(payload) { + return mssqlRawError{}, pos, fmt.Errorf("mssql: truncated error token") + } + size := int(binary.LittleEndian.Uint16(payload[pos : pos+2])) + end := pos + 2 + size + if size < 6 || end > len(payload) || pos+8 > len(payload) { + return mssqlRawError{}, pos, fmt.Errorf("mssql: invalid error token size") + } + pos += 2 + number := int32(binary.LittleEndian.Uint32(payload[pos : pos+4])) + pos += 4 + pos += 2 + message, next, err := mssqlReadUSVarChar(payload, pos) + if err != nil { + return mssqlRawError{}, pos, err + } + return mssqlRawError{number: number, message: message}, end, mssqlEnsureSkipBVarStrings(payload, next, end) +} + +func mssqlSkipUSVarError(payload []byte, pos int) (int, error) { + if pos+2 > len(payload) { + return pos, fmt.Errorf("mssql: truncated info token") + } + size := int(binary.LittleEndian.Uint16(payload[pos : pos+2])) + end := pos + 2 + size + if size < 6 || end > len(payload) || pos+8 > len(payload) { + return pos, fmt.Errorf("mssql: invalid info token size") + } + _, _, err := mssqlReadUSVarChar(payload, pos+8) + return end, err +} + +func mssqlEnsureSkipBVarStrings(payload []byte, pos, end int) error { + for i := 0; i < 2; i++ { + if pos >= end { + return fmt.Errorf("mssql: truncated string in error token") + } + length := int(payload[pos]) * 2 + pos++ + if pos+length > end { + return fmt.Errorf("mssql: invalid string in error token") + } + pos += length + } + if pos+4 > end { + return fmt.Errorf("mssql: truncated error line number") + } + return nil +} + +func mssqlSkipLen16(payload []byte, pos int) (int, error) { + if pos+2 > len(payload) { + return pos, fmt.Errorf("mssql: truncated token") + } + size := int(binary.LittleEndian.Uint16(payload[pos : pos+2])) + next := pos + 2 + size + if next > len(payload) { + return pos, fmt.Errorf("mssql: invalid token size") + } + return next, nil +} + +func mssqlReadUSVarChar(payload []byte, pos int) (string, int, error) { + if pos+2 > len(payload) { + return "", pos, fmt.Errorf("mssql: truncated us varchar") + } + chars := int(binary.LittleEndian.Uint16(payload[pos : pos+2])) + pos += 2 + size := chars * 2 + if pos+size > len(payload) { + return "", pos, fmt.Errorf("mssql: invalid us varchar size") + } + return mssqlDecodeUCS2(payload[pos : pos+size]), pos + size, nil +} + +func mssqlWritePacket(w io.Writer, packetType byte, payload []byte) error { + if len(payload)+8 > 0xffff { + return fmt.Errorf("mssql: packet too large") + } + header := []byte{packetType, tdsStatusEOM, 0, 0, 0, 0, 1, 0} + binary.BigEndian.PutUint16(header[2:4], uint16(len(payload)+8)) + if _, err := w.Write(header); err != nil { + return err + } + _, err := w.Write(payload) + return err +} + +func mssqlReadMessage(r io.Reader) (byte, []byte, error) { + var packetType byte + var payload []byte + for { + header := make([]byte, 8) + if _, err := io.ReadFull(r, header); err != nil { + return 0, nil, err + } + if packetType == 0 { + packetType = header[0] + } else if packetType != header[0] { + return 0, nil, fmt.Errorf("mssql: packet type changed in message") + } + size := int(binary.BigEndian.Uint16(header[2:4])) + if size < 8 { + return 0, nil, fmt.Errorf("mssql: invalid packet size") + } + chunk := make([]byte, size-8) + if _, err := io.ReadFull(r, chunk); err != nil { + return 0, nil, err + } + payload = append(payload, chunk...) + if header[1]&tdsStatusEOM != 0 { + return packetType, payload, nil + } + } +} + +func mssqlUCS2(s string) []byte { + runes := utf16.Encode([]rune(s)) + out := make([]byte, len(runes)*2) + for i, r := range runes { + binary.LittleEndian.PutUint16(out[i*2:], r) + } + return out +} + +func mssqlDecodeUCS2(data []byte) string { + if len(data)%2 != 0 { + data = data[:len(data)-1] + } + runes := make([]uint16, len(data)/2) + for i := range runes { + runes[i] = binary.LittleEndian.Uint16(data[i*2:]) + } + return string(utf16.Decode(runes)) +} + +func mssqlEncodePassword(password string) []byte { + out := mssqlUCS2(password) + for i, ch := range out { + out[i] = (((ch << 4) & 0xff) | (ch >> 4)) ^ 0xa5 + } + return out +}