25 Commits
Author SHA1 Message Date
claudemiro 6c55af2b5c Adding README and LICENSE to distribution 2016-03-05 09:43:02 -03:00
claudemiro 679ff1e589 Adding the version in the generated binary 2016-03-05 09:30:36 -03:00
claudemiro cbdd4e428a Better config example 2016-03-05 09:29:57 -03:00
claudemiro 805e9e3957 Removed Makefile and Add a Rakefile 2016-03-05 09:27:40 -03:00
claudemiro 44bbc1b1ea Renamed constants, avoid export uncessary constants 2016-03-04 08:33:18 -03:00
claudemiro e02773bbeb function init is not necessary 2016-03-04 08:32:39 -03:00
claudemiro 5fa621d743 Added user-agent header with the url of the project 2016-03-04 08:32:08 -03:00
claudemiro 880d7bda2a Using reader for reading config file 2016-03-04 08:31:32 -03:00
claudemiro 7dad2de275 Added SSL port 2016-03-04 08:31:03 -03:00
claudemiro 081611fd87 Changed the banner, fixed spelling 2016-03-04 08:30:26 -03:00
claudemiro 8810e523f0 Adding config for functional tests 2016-03-03 21:39:15 -03:00
claudemiro baf551ecdf Updated TODO.org 2016-03-03 21:07:57 -03:00
claudemiro f3dba811b7 Improved Handling SSL
* Panic is betten than using channels in this case
* Changed the names to a more recognized names
* Improved the README with a full config file example
2016-03-03 21:01:21 -03:00
claudemiro 55a7a0b96c Fixed SSL verification 2016-03-03 20:46:37 -03:00
Claudemiro cadee4e26d Merge pull request #10 from zaeznet/SSL
SSL Implementation
2016-03-03 20:45:42 -03:00
Welington Sampaio 03e94887a7 Both server running. SSL and Non SSL 2016-03-03 17:12:51 -03:00
Welington Sampaio e93145efe7 Both server running. SSL and Non SSL 2016-03-03 14:54:53 -03:00
Welington Sampaio 051da185e5 SSL Implementation 2016-03-03 08:37:00 -03:00
claudemiro 51eacdcf51 Managing vendor dependencies with glide 2016-02-21 20:15:10 -03:00
Claudemiro 412d8be451 Merge pull request #7 from jweslley/private-channel-payload
Private channel payload
2016-01-28 07:41:57 -02:00
Claudemiro 68f2e15320 Merge pull request #6 from jweslley/master
minor improvements
2016-01-27 09:10:52 -02:00
Jonhnny Weslley a77fb791f7 sign private channels with additional user data payload
Pusher accepts js clients to use the same auth endpoint to both presence
and private channels.
2016-01-24 12:43:06 -03:00
Jonhnny Weslley 8bf0786cb5 add systemd config file 2016-01-24 12:06:38 -03:00
Jonhnny Weslley ba3b699517 refactor common logic 2016-01-24 11:54:06 -03:00
claudemiro 517462a7af Cache of the valid channel name regex
* Added a Benchmark to verify the performance
* Now it is 10 times faster
2016-01-23 22:25:14 -02:00
70 changed files with 2298 additions and 353 deletions
+2
View File
@@ -155,4 +155,6 @@ flymake*
ignore_http/* ignore_http/*
config.json config.json
*.pem
build
-13
View File
@@ -1,13 +0,0 @@
default: debug
debug:
GO15VENDOREXPERIMENT=1 go install -ldflags "-w" github.com/dimiro1/ipe
run-debug: debug
${GOPATH}/bin/ipe --config ${GOPATH}/src/github.com/dimiro1/ipe/config.json -logtostderr=true -v=2
test:
GO15VENDOREXPERIMENT=1 go test `go list ./... | grep -v vendor`
dev-deps:
go get github.com/pusher/pusher-http-go
+17 -5
View File
@@ -25,6 +25,12 @@ This software is written in Go - the WYSIWYG lang
* Multiple apps in the same instance; * Multiple apps in the same instance;
* Drop in replacement for pusher server; * Drop in replacement for pusher server;
# Download pre built binaries
You can download pre built binaries from the [releases tab](https://github.com/dimiro1/ipe/releases).
I do not have a Windows machine, so I can only distribute binaries for amd64 linux and amd64 darwin.
# Building # Building
```console ```console
@@ -44,19 +50,25 @@ $ go install github.com/dimiro1/ipe
```json ```json
{ {
"Host": ":8080", "Host": ":8080",
"SSL": false,
"SSLHost": ":4433",
"SSLKeyFile": "A key.pem file",
"SSLCertFile": "A cert.pem file",
"Apps": [ "Apps": [
{ {
"ApplicationDisabled": false, "ApplicationDisabled": false,
"Secret": "APP_SECRET", "Secret": "A really secret random string",
"Key": "APP_KEY", "Key": "A random Key string",
"Name": "APP_NAME", "OnlySSL": false,
"AppID": "APP_ID", "Name": "The app name",
"AppID": "The app ID",
"UserEvents": true, "UserEvents": true,
"WebHooks": true, "WebHooks": true,
"URLWebHook": "http://localhost:4567/php/hook.php" "URLWebHook": "Some URL to send webhooks"
} }
] ]
} }
``` ```
## Libraries ## Libraries
+58
View File
@@ -0,0 +1,58 @@
# Copyright 2016 Claudemiro Alves Feitosa Neto. All rights reserved.
# Use of this source code is governed by a MIT-style
# license that can be found in the LICENSE file.
require 'rake/clean'
VERSION = 'v1.0.0'
GITHASH = `git rev-parse --short HEAD`
DATE = Time.now.strftime '%Y%m%d%H%M%S'
CLOBBER.include 'build'
task :default => [:'run-debug']
desc 'Build a debug version'
task :debug do
sh "GO15VENDOREXPERIMENT=1 go install -ldflags '-w -X main.version=DEBUG -X main.buildstamp=DEBUG -X main.githash=DEBUG' github.com/dimiro1/ipe"
end
desc 'Build and run debug version'
task :'run-debug' => :debug do
sh '$GOPATH/bin/ipe --config $GOPATH/src/github.com/dimiro1/ipe/config.json -logtostderr=true -v=2'
end
desc 'Run test suite'
task :test do
sh 'GO15VENDOREXPERIMENT=1 go test . `glide nv`'
end
desc 'Download the development dependencies'
task :'dev-deps' do
sh 'go get github.com/pusher/pusher-http-go'
end
desc 'Generate distributions'
task :distribute => [:linux, :darwin]
desc 'Generate a linux distribution'
task :linux do
Rake::Task['build'].invoke 'linux'
end
desc 'Generate a darwin distribution'
task :darwin do
Rake::Task['build'].invoke 'darwin'
end
task :build, [:os] do |t, args|
t.reenable
os = args[:os]
sh "mkdir -p build/#{os}"
sh "GO15VENDOREXPERIMENT=1 GOOS=#{os} GOARCH=amd64 go build -ldflags '-X main.version=#{VERSION} -X main.buildstamp=#{DATE} -X main.githash=#{GITHASH}' -o build/#{os}/ipe github.com/dimiro1/ipe"
sh "cp ipe/config-example.json build/#{os}/config.json"
sh "cp LICENSE build/#{os}/"
sh "cp README.md build/#{os}/"
sh "tar -C build/#{os} -czf build/ipe_#{VERSION}_#{os}_amd64.tar.gz ."
end
+2 -2
View File
@@ -1,12 +1,12 @@
IPÊ IPÊ
--- ---
* TODO [11/14] * TODO [12/14]
* [X] Autenticação API Rest * [X] Autenticação API Rest
* [X] Autenticação Websockets * [X] Autenticação Websockets
* [X] Ping e Pong * [X] Ping e Pong
* [ ] Escrever testes automatizados * [ ] Escrever testes automatizados
* [ ] SSL * [X] SSL
* [X] Expvar - Canais, inscritos * [X] Expvar - Canais, inscritos
* [X] Otimizações [3/3] * [X] Otimizações [3/3]
* [X] Refatorar partes do código, remover repetições * [X] Refatorar partes do código, remover repetições
+1 -1
View File
@@ -1,2 +1,2 @@
client: go run client.go client: go run client.go
server: go run ../main.go -config ./config.json -logtostderr server: go run ../main.go -config ./functional-config.json -logtostderr
+20
View File
@@ -0,0 +1,20 @@
{
"Host": ":8080",
"Encrypted": false,
"SSLHost": ":8090",
"SSLKeyFile": "key.pem",
"SSLCertFile": "cert.pem",
"Apps": [
{
"ApplicationDisabled": false,
"OnlySSL": false,
"Secret": "7ad3753142a6693b25b9",
"Key": "278d525bdf162c739803",
"Name": "App for Functional Test",
"AppID": "1",
"UserEvents": true,
"WebHooks": false,
"URLWebHook": "http://127.0.0.1:4567/php/hook.php"
}
]
}
Generated
+14
View File
@@ -0,0 +1,14 @@
hash: bfb508bdf4f85c71c7eef4bb9431148ee373dfaf31d001d050d933f80ac4288e
updated: 2016-02-21T20:13:57.03748112-03:00
imports:
- name: github.com/golang/glog
version: 23def4e6c14b4da8ac2ed8007337bc5eb5007998
- name: github.com/gorilla/context
version: 1c83b3eabd45b6d76072b66b746c20815fb2872d
- name: github.com/gorilla/mux
version: 26a6070f849969ba72b72256e9f14cf519751690
- name: github.com/gorilla/websocket
version: 5c91b59efa232fa9a87b705d54101832c498a172
- name: github.com/pusher/pusher-http-go
version: 8d4ffe157699620440932e4d03253a22533f2e43
devImports: []
+6
View File
@@ -0,0 +1,6 @@
package: github.com/dimiro1/ipe
import:
- package: github.com/golang/glog
- package: github.com/gorilla/mux
- package: github.com/gorilla/websocket
- package: github.com/pusher/pusher-http-go
+14
View File
@@ -0,0 +1,14 @@
[Unit]
Description=Ipe
After=syslog.target network.target
[Service]
Type=simple
User=ipe
StandardOutput=syslog
StandardError=syslog
SyslogIdentifier=ipe
ExecStart=/path/to/ipe -logtostderr -config="path_to_config.json"
[Install]
WantedBy=multi-user.target
+3 -2
View File
@@ -11,6 +11,7 @@ import (
"sync" "sync"
"time" "time"
"github.com/dimiro1/ipe/utils"
log "github.com/golang/glog" log "github.com/golang/glog"
) )
@@ -40,12 +41,12 @@ func (c *channel) IsPublic() bool {
// Check if the type of the channel is presence // Check if the type of the channel is presence
func (c *channel) IsPresence() bool { func (c *channel) IsPresence() bool {
return strings.HasPrefix(c.ChannelID, "presence-") return utils.IsPresenceChannel(c.ChannelID)
} }
// Check if the type of the channel is private // Check if the type of the channel is private
func (c *channel) IsPrivate() bool { func (c *channel) IsPrivate() bool {
return strings.HasPrefix(c.ChannelID, "private-") return utils.IsPrivateChannel(c.ChannelID)
} }
// Get the total of subscribers // Get the total of subscribers
+10 -15
View File
@@ -1,25 +1,20 @@
{ {
"Host": ":8080", "Host": ":8080",
"SSL": false,
"SSLHost": ":4433",
"SSLKeyFile": "A key.pem file",
"SSLCertFile": "A cert.pem file",
"Apps": [ "Apps": [
{ {
"ApplicationDisabled": false, "ApplicationDisabled": false,
"Secret": "7ad3753142a6693b25b9", "Secret": "A really secret random string",
"Key": "278d525bdf162c739803", "Key": "A random Key string",
"Name": "App 1", "OnlySSL": false,
"AppID": "321", "Name": "The app name",
"AppID": "The app ID",
"UserEvents": true, "UserEvents": true,
"WebHooks": true, "WebHooks": true,
"URLWebHook": "http://127.0.0.1:4567/php/hook.php" "URLWebHook": "Some URL to send webhooks"
},
{
"ApplicationDisabled": false,
"Secret": "d6824d2fa32888931504",
"Key": "c8b30f611ffb13202976",
"Name": "App 2",
"AppID": "123",
"UserEvents": true,
"WebHooks": false,
"URLWebHook": "http://127.0.0.1:4567/php/hook.php"
} }
] ]
} }
+9 -4
View File
@@ -11,10 +11,15 @@ import (
// The config file // The config file
type configFile struct { type configFile struct {
Host string // The host, eg: :8080 will start on 0.0.0.0:8080 Host string // The host, eg: :8080 will start on 0.0.0.0:8080
User string User string
Password string Password string
Apps []*app SSL bool
SSLHost string
SSLKeyFile string
SSLCertFile string
Apps []*app
} }
// Initialize Apps // Initialize Apps
+16 -16
View File
@@ -9,39 +9,39 @@ const (
// 4000 - 4099 // 4000 - 4099
// Indicates an error resulting in the connection being closed by Pusher, // Indicates an error resulting in the connection being closed by Pusher,
// and that attempting to reconnect using the same parameters will not succeed. // and that attempting to reconnect using the same parameters will not succeed.
APPLICATION_ONLY_ACCEPTS_SSL = 4000 applicationOnlyAcceptsSSL = 4000
APPLICATION_DOES_NOT_EXISTS = 4001 applicationDoesNotExists = 4001
APPLICATION_DISABLED = 4003 applicationDisabled = 4003
APPLICATION_IS_OVER_CONNECTION_QUOTA = 4004 // Not Implemented applicationIsOverConnectionQuota = 4004 // Not Implemented
PATH_NOT_FOUND = 4005 // Not Implemented pathNotFound = 4005 // Not Implemented
INVALID_VERSION_STRING_FORMAT = 4006 invalidVersionStringFormat = 4006
UNSUPPORTED_PROTOCOL_VERSION = 4007 unsupportedProtocolVersion = 4007
NO_PROTOCOL_VERSION_SUPPLIED = 4008 noProtocolVersionSupplied = 4008
// 4100 - 4199 // 4100 - 4199
// Indicates an error resulting in the connection being closed by Pusher, // Indicates an error resulting in the connection being closed by Pusher,
// and the client may reconnect after 1s or more // and the client may reconnect after 1s or more
OVER_CAPACITY = 4100 // Not Implemented overCapacity = 4100 // Not Implemented
// 4200 - 4299 // 4200 - 4299
// Indicate an error resulting in the connection being closed by Pusher, // Indicate an error resulting in the connection being closed by Pusher,
// and the client my reconnect immediately // and the client my reconnect immediately
GENERIC_RECONNECT_IMMEDIATELY = 4200 genericReconnectImmediately = 4200
PONG_REPLY_NOT_RECEIVED = 4201 // Ping was sent to the client, but no reply was received; Not Implemented pongReplyNotReceived = 4201 // Ping was sent to the client, but no reply was received; Not Implemented
CLOSED_AFTER_INACTIVITY = 4202 // Client has been inactive for a long time (24 hours) and client does not suppot ping.; Not Implemented closedAfterInactivity = 4202 // Client has been inactive for a long time (24 hours) and client does not suppot ping.; Not Implemented
// 4300 - 4399 // 4300 - 4399
// Any other type of error // Any other type of error
CLIENT_REJECTED_DUE_TO_RATE_LIMIT = 4301 // Not Implemented clientRejectedDueToRateLimit = 4301 // Not Implemented
// Pusher send null, This app use this error code to send the null value // Pusher send null, This app use this error code to send the null value
// see ErrorEvent // see ErrorEvent
GENERIC_ERROR = 0 otherError = 0
) )
// Only this version is supported // Only this version is supported
const SUPPORTED_PROTOCOL_VERSION = 7 const supportedProtocolVersion = 7
// // Maximun event size permitted 10 kB // // Maximun event size permitted 10 kB
// See: http://blogs.gnome.org/cneumair/2008/09/30/1-kb-1024-bytes-no-1-kb-1000-bytes/ // See: http://blogs.gnome.org/cneumair/2008/09/30/1-kb-1024-bytes-no-1-kb-1000-bytes/
const MAX_DATA_EVENT_SIZE = 10 * 1000 const maxDataEventSize = 10 * 1000
+8 -8
View File
@@ -31,7 +31,7 @@ type unsupportedProtocolVersionError struct {
func newUnsupportedProtocolVersionError() unsupportedProtocolVersionError { func newUnsupportedProtocolVersionError() unsupportedProtocolVersionError {
return unsupportedProtocolVersionError{ return unsupportedProtocolVersionError{
baseWebsocketError{Code: UNSUPPORTED_PROTOCOL_VERSION, Msg: "Unsupported protocol version"}, baseWebsocketError{Code: unsupportedProtocolVersion, Msg: "Unsupported protocol version"},
} }
} }
@@ -43,7 +43,7 @@ type applicationDoesNotExistsError struct {
func newApplicationDoesNotExistsError() applicationDoesNotExistsError { func newApplicationDoesNotExistsError() applicationDoesNotExistsError {
return applicationDoesNotExistsError{ return applicationDoesNotExistsError{
baseWebsocketError{Code: APPLICATION_DOES_NOT_EXISTS, Msg: "Could not found an app with the given key"}, baseWebsocketError{Code: applicationDoesNotExists, Msg: "Could not found an app with the given key"},
} }
} }
@@ -54,7 +54,7 @@ type noProtocolVersionSuppliedError struct {
func newNoProtocolVersionSuppliedError() noProtocolVersionSuppliedError { func newNoProtocolVersionSuppliedError() noProtocolVersionSuppliedError {
return noProtocolVersionSuppliedError{ return noProtocolVersionSuppliedError{
baseWebsocketError{Code: NO_PROTOCOL_VERSION_SUPPLIED, Msg: "No protocol version supplied"}, baseWebsocketError{Code: noProtocolVersionSupplied, Msg: "No protocol version supplied"},
} }
} }
@@ -66,7 +66,7 @@ type applicationDisabledError struct {
func newApplicationDisabledError() noProtocolVersionSuppliedError { func newApplicationDisabledError() noProtocolVersionSuppliedError {
return noProtocolVersionSuppliedError{ return noProtocolVersionSuppliedError{
baseWebsocketError{Code: APPLICATION_DISABLED, Msg: "Application disabled"}, baseWebsocketError{Code: applicationDisabled, Msg: "Application disabled"},
} }
} }
@@ -77,7 +77,7 @@ type applicationOnlyAccepsSSLError struct {
func newApplicationOnlyAccepsSSLError() applicationOnlyAccepsSSLError { func newApplicationOnlyAccepsSSLError() applicationOnlyAccepsSSLError {
return applicationOnlyAccepsSSLError{ return applicationOnlyAccepsSSLError{
baseWebsocketError{Code: APPLICATION_ONLY_ACCEPTS_SSL, Msg: "Application only accepts SSL connections, reconnect using wss://"}, baseWebsocketError{Code: applicationOnlyAcceptsSSL, Msg: "Application only accepts SSL connections, reconnect using wss://"},
} }
} }
@@ -88,7 +88,7 @@ type invalidVersionStringFormatError struct {
func newInvalidVersionStringFormatError() invalidVersionStringFormatError { func newInvalidVersionStringFormatError() invalidVersionStringFormatError {
return invalidVersionStringFormatError{ return invalidVersionStringFormatError{
baseWebsocketError{Code: INVALID_VERSION_STRING_FORMAT, Msg: "Invalid version string format"}, baseWebsocketError{Code: invalidVersionStringFormat, Msg: "Invalid version string format"},
} }
} }
@@ -101,7 +101,7 @@ type genericReconnectImmediatelyError struct {
func newGenericReconnectImmediatelyError() genericReconnectImmediatelyError { func newGenericReconnectImmediatelyError() genericReconnectImmediatelyError {
return genericReconnectImmediatelyError{ return genericReconnectImmediatelyError{
baseWebsocketError{Code: GENERIC_RECONNECT_IMMEDIATELY, Msg: "Generic reconnect immediately"}, baseWebsocketError{Code: genericReconnectImmediately, Msg: "Generic reconnect immediately"},
} }
} }
@@ -113,6 +113,6 @@ type genericError struct {
func newGenericError(msg string) genericError { func newGenericError(msg string) genericError {
return genericError{ return genericError{
baseWebsocketError{Code: GENERIC_ERROR, Msg: msg}, baseWebsocketError{Code: otherError, Msg: msg},
} }
} }
+1 -1
View File
@@ -164,7 +164,7 @@ type errorEvent struct {
func newErrorEvent(code int, message string) errorEvent { func newErrorEvent(code int, message string) errorEvent {
var data interface{} var data interface{}
if code == GENERIC_ERROR { if code == otherError {
data = struct { data = struct {
Code *int `json:"code"` Code *int `json:"code"`
Message string `json:"message"` Message string `json:"message"`
+17 -10
View File
@@ -6,34 +6,41 @@ package ipe
import ( import (
"encoding/json" "encoding/json"
"io/ioutil"
"net/http"
"math/rand" "math/rand"
"net/http"
"os"
"time" "time"
log "github.com/golang/glog"
) )
// Conf holds the global configuration state // Conf holds the global configuration state
var conf configFile var conf configFile
// Start Parse the configuration file and starts the ipe server // Start Parse the configuration file and starts the ipe server
func Start(configfile string) error { // It Panic if could not start the HTTP or HTTPS server
func Start(configfile string) {
rand.Seed(time.Now().Unix()) rand.Seed(time.Now().Unix())
file, err := ioutil.ReadFile(configfile) file, err := os.Open(configfile)
if err != nil { if err != nil {
return err log.Fatal(err)
} }
if err := json.Unmarshal(file, &conf); err != nil { if err := json.NewDecoder(file).Decode(&conf); err != nil {
return err log.Fatal(err)
} }
conf.Init() conf.Init()
router := newRouter() router := newRouter()
if err := http.ListenAndServe(conf.Host, router); err != nil { if conf.SSL {
return err go func() {
log.Infof("Starting HTTPS service on %s ...", conf.SSLHost)
log.Fatal(http.ListenAndServeTLS(conf.SSLHost, conf.SSLCertFile, conf.SSLKeyFile, router))
}()
} }
return nil log.Infof("Starting HTTP service on %s ...", conf.Host)
log.Fatal(http.ListenAndServe(conf.Host, router))
} }
+3 -2
View File
@@ -10,6 +10,7 @@ import (
"net/http" "net/http"
"strings" "strings"
"github.com/dimiro1/ipe/utils"
log "github.com/golang/glog" log "github.com/golang/glog"
"github.com/gorilla/mux" "github.com/gorilla/mux"
) )
@@ -55,7 +56,7 @@ func postEvents(w http.ResponseWriter, r *http.Request) {
} }
// The event data should not be larger than 10KB. // The event data should not be larger than 10KB.
if len(input.Data) > MAX_DATA_EVENT_SIZE { if len(input.Data) > maxDataEventSize {
http.Error(w, "Request too large.", http.StatusRequestEntityTooLarge) http.Error(w, "Request too large.", http.StatusRequestEntityTooLarge)
return return
} }
@@ -272,7 +273,7 @@ func getChannelUsers(w http.ResponseWriter, r *http.Request) {
appID := vars["app_id"] appID := vars["app_id"]
channelName := vars["channel_name"] channelName := vars["channel_name"]
isPresence := strings.HasPrefix(channelName, "presence-") isPresence := utils.IsPresenceChannel(channelName)
if !isPresence { if !isPresence {
http.Error(w, "This api endpoint is restricted to presence channels.", http.StatusBadRequest) http.Error(w, "This api endpoint is restricted to presence channels.", http.StatusBadRequest)
+1
View File
@@ -150,6 +150,7 @@ func triggerHook(name string, a *app, c *channel, event hookEvent) {
return return
} }
req.Header.Set("User-Agent", "Ipe UA; (+https://github.com/dimiro1/ipe)")
req.Header.Set("Content-Type", "application/json") req.Header.Set("Content-Type", "application/json")
req.Header.Set("X-Pusher-Key", a.Key) req.Header.Set("X-Pusher-Key", a.Key)
req.Header.Set("X-Pusher-Signature", utils.HashMAC(js, []byte(a.Secret))) req.Header.Set("X-Pusher-Signature", utils.HashMAC(js, []byte(a.Secret)))
+6 -6
View File
@@ -39,12 +39,12 @@ func onOpen(conn *websocket.Conn, w http.ResponseWriter, r *http.Request, sessio
switch { switch {
case strings.TrimSpace(p) == "": case strings.TrimSpace(p) == "":
return newNoProtocolVersionSuppliedError() return newNoProtocolVersionSuppliedError()
case protocol != SUPPORTED_PROTOCOL_VERSION: case protocol != supportedProtocolVersion:
return newUnsupportedProtocolVersionError() return newUnsupportedProtocolVersionError()
case app.ApplicationDisabled: case app.ApplicationDisabled:
return newApplicationDisabledError() return newApplicationDisabledError()
case r.TLS != nil: case app.OnlySSL:
if app.OnlySSL { if r.TLS == nil {
return newApplicationOnlyAccepsSSLError() return newApplicationOnlyAccepsSSLError()
} }
} }
@@ -123,13 +123,13 @@ func onMessage(conn *websocket.Conn, w http.ResponseWriter, r *http.Request, ses
break break
} }
isPresence := strings.HasPrefix(channelName, "presence-") isPresence := utils.IsPresenceChannel(channelName)
isPrivate := strings.HasPrefix(channelName, "private-") isPrivate := utils.IsPrivateChannel(channelName)
if isPresence || isPrivate { if isPresence || isPrivate {
toSign := []string{connection.SocketID, channelName} toSign := []string{connection.SocketID, channelName}
if isPresence { if isPresence || len(subscribeEvent.Data.ChannelData) > 0 {
toSign = append(toSign, subscribeEvent.Data.ChannelData) toSign = append(toSign, subscribeEvent.Data.ChannelData)
} }
+27 -15
View File
@@ -9,32 +9,44 @@ import (
"fmt" "fmt"
"github.com/dimiro1/ipe/ipe" "github.com/dimiro1/ipe/ipe"
log "github.com/golang/glog"
) )
// Main function, initialze the system // These variables are generated by the linker
// please see the makefile for mor information.
var (
version string = "version"
buildstamp string = "buildstamp"
githash string = "githash"
)
// Main function, initialize the system
func main() { func main() {
var filename = flag.String("config", "config.json", "Config file location") var filename = flag.String("config", "config.json", "Config file location")
flag.Parse() flag.Parse()
printBanner() printBanner()
if err := ipe.Start(*filename); err != nil { ipe.Start(*filename)
log.Fatal(err)
}
} }
// Print a beautifull banner // Print a beautiful banner
func printBanner() { func printBanner() {
fmt.Print("\033[36m") fmt.Print("\033[31m")
fmt.Print(` fmt.Print(`
██╗██████╗ ███████╗ d8b
██║██╔══██╗██╔════╝ Y8P
██║██████╔╝█████╗
██║██╔═══╝ ██╔══╝ 888 88888b. .d88b.
██║██║ ███████╗ 888 888 "88b d8P Y8b
╚═╝╚═╝ ╚══════╝`) 888 888 888 88888888
888 888 d88P Y8b.
888 88888P" "Y8888
888
888
888
`)
fmt.Println("\033[0m") fmt.Println("\033[0m")
fmt.Println("\033[32mWelcome to Ipê - Yet another Pusher server clone\033[0m") fmt.Println("\033[32mWelcome to Ipê - Yet another Pusher server clone (https://github.com/dimiro1/ipe)\033[0m")
fmt.Printf("\033[32mVersion %s+%s.git.%s\033[0m\n", version, buildstamp, githash)
fmt.Println("\033[33mBy: Claudemiro Alves Feitosa Neto <[email protected]>\033[0m") fmt.Println("\033[33mBy: Claudemiro Alves Feitosa Neto <[email protected]>\033[0m")
} }
+11 -8
View File
@@ -12,8 +12,11 @@ import (
"math" "math"
"math/rand" "math/rand"
"regexp" "regexp"
"strings"
) )
var validChannelName *regexp.Regexp = regexp.MustCompile("^[A-Za-z0-9_\\-=@,.;]+$")
// HashMAC Calculates the MAC signing with the given key and returns the hexadecimal encoded Result // HashMAC Calculates the MAC signing with the given key and returns the hexadecimal encoded Result
func HashMAC(message, key []byte) string { func HashMAC(message, key []byte) string {
mac := hmac.New(sha256.New, key) mac := hmac.New(sha256.New, key)
@@ -25,18 +28,18 @@ func HashMAC(message, key []byte) string {
// GenerateSessionID Generate a new random Hash // GenerateSessionID Generate a new random Hash
func GenerateSessionID() string { func GenerateSessionID() string {
MAX := math.MaxInt64 return fmt.Sprintf("%d.%d", rand.Intn(math.MaxInt64), rand.Intn(math.MaxInt64))
return fmt.Sprintf("%d.%d", rand.Intn(MAX), rand.Intn(MAX))
} }
// IsChannelNameValid Verify if the channel name is valid // IsChannelNameValid Verify if the channel name is valid
func IsChannelNameValid(channelName string) bool { func IsChannelNameValid(channelName string) bool {
matched, err := regexp.MatchString("^[A-Za-z0-9_\\-=@,.;]+$", channelName) return validChannelName.Match([]byte(channelName))
}
if err == nil && matched { func IsPrivateChannel(channelName string) bool {
return true return strings.HasPrefix(channelName, "private-")
} }
return false func IsPresenceChannel(channelName string) bool {
return strings.HasPrefix(channelName, "presence-")
} }
+12
View File
@@ -9,6 +9,18 @@ import (
"testing" "testing"
) )
func BenchmarkGenerateSession(b *testing.B) {
for i := 0; i < b.N; i++ {
GenerateSessionID()
}
}
func BenchmarkIsChannelNameValid(b *testing.B) {
for i := 0; i < b.N; i++ {
IsChannelNameValid("hello-world")
}
}
func TestGenerateSession(t *testing.T) { func TestGenerateSession(t *testing.T) {
sessionID := GenerateSessionID() sessionID := GenerateSessionID()
+1 -1
View File
@@ -5,7 +5,7 @@ Leveled execution logs for Go.
This is an efficient pure Go implementation of leveled logs in the This is an efficient pure Go implementation of leveled logs in the
manner of the open source C++ package manner of the open source C++ package
http://code.google.com/p/google-glog https://github.com/google/glog
By binding methods to booleans it is possible to use the log package By binding methods to booleans it is possible to use the log package
without paying the expense of evaluating the arguments to the log. without paying the expense of evaluating the arguments to the log.
+4 -1
View File
@@ -676,7 +676,10 @@ func (l *loggingT) output(s severity, buf *buffer, file string, line int, alsoTo
} }
} }
data := buf.Bytes() data := buf.Bytes()
if l.toStderr { if !flag.Parsed() {
os.Stderr.Write([]byte("ERROR: logging before flag.Parse: "))
os.Stderr.Write(data)
} else if l.toStderr {
os.Stderr.Write(data) os.Stderr.Write(data)
} else { } else {
if alsoToStderr || l.alsoToStderr || s >= l.stderrThreshold.get() { if alsoToStderr || l.alsoToStderr || s >= l.stderrThreshold.get() {
+2 -3
View File
@@ -1,9 +1,8 @@
language: go language: go
sudo: false
go: go:
- 1.0
- 1.1
- 1.2
- 1.3 - 1.3
- 1.4 - 1.4
- 1.5
- tip - tip
+11 -4
View File
@@ -1,7 +1,14 @@
language: go language: go
sudo: false
go: go:
- 1.0 - 1.3
- 1.1 - 1.4
- 1.2 - 1.5
- tip - tip
install:
- go get golang.org/x/tools/cmd/vet
script:
- go get -t -v ./...
- diff -u <(echo -n) <(gofmt -d -s .)
- go tool vet .
- go test -v -race ./...
Generated Vendored Executable → Regular
View File
Generated Vendored Executable → Regular
+236 -3
View File
@@ -1,7 +1,240 @@
mux mux
=== ===
[![Build Status](https://travis-ci.org/gorilla/mux.png?branch=master)](https://travis-ci.org/gorilla/mux) [![GoDoc](https://godoc.org/github.com/gorilla/mux?status.svg)](https://godoc.org/github.com/gorilla/mux)
[![Build Status](https://travis-ci.org/gorilla/mux.svg?branch=master)](https://travis-ci.org/gorilla/mux)
gorilla/mux is a powerful URL router and dispatcher. Package `gorilla/mux` implements a request router and dispatcher.
Read the full documentation here: http://www.gorillatoolkit.org/pkg/mux The name mux stands for "HTTP request multiplexer". Like the standard `http.ServeMux`, `mux.Router` matches incoming requests against a list of registered routes and calls a handler for the route that matches the URL or other conditions. The main features are:
* Requests can be matched based on URL host, path, path prefix, schemes, header and query values, HTTP methods or using custom matchers.
* URL hosts and paths can have variables with an optional regular expression.
* Registered URLs can be built, or "reversed", which helps maintaining references to resources.
* Routes can be used as subrouters: nested routes are only tested if the parent route matches. This is useful to define groups of routes that share common conditions like a host, a path prefix or other repeated attributes. As a bonus, this optimizes request matching.
* It implements the `http.Handler` interface so it is compatible with the standard `http.ServeMux`.
Let's start registering a couple of URL paths and handlers:
```go
func main() {
r := mux.NewRouter()
r.HandleFunc("/", HomeHandler)
r.HandleFunc("/products", ProductsHandler)
r.HandleFunc("/articles", ArticlesHandler)
http.Handle("/", r)
}
```
Here we register three routes mapping URL paths to handlers. This is equivalent to how `http.HandleFunc()` works: if an incoming request URL matches one of the paths, the corresponding handler is called passing (`http.ResponseWriter`, `*http.Request`) as parameters.
Paths can have variables. They are defined using the format `{name}` or `{name:pattern}`. If a regular expression pattern is not defined, the matched variable will be anything until the next slash. For example:
```go
r := mux.NewRouter()
r.HandleFunc("/products/{key}", ProductHandler)
r.HandleFunc("/articles/{category}/", ArticlesCategoryHandler)
r.HandleFunc("/articles/{category}/{id:[0-9]+}", ArticleHandler)
```
The names are used to create a map of route variables which can be retrieved calling `mux.Vars()`:
```go
vars := mux.Vars(request)
category := vars["category"]
```
And this is all you need to know about the basic usage. More advanced options are explained below.
Routes can also be restricted to a domain or subdomain. Just define a host pattern to be matched. They can also have variables:
```go
r := mux.NewRouter()
// Only matches if domain is "www.example.com".
r.Host("www.example.com")
// Matches a dynamic subdomain.
r.Host("{subdomain:[a-z]+}.domain.com")
```
There are several other matchers that can be added. To match path prefixes:
```go
r.PathPrefix("/products/")
```
...or HTTP methods:
```go
r.Methods("GET", "POST")
```
...or URL schemes:
```go
r.Schemes("https")
```
...or header values:
```go
r.Headers("X-Requested-With", "XMLHttpRequest")
```
...or query values:
```go
r.Queries("key", "value")
```
...or to use a custom matcher function:
```go
r.MatcherFunc(func(r *http.Request, rm *RouteMatch) bool {
return r.ProtoMajor == 0
})
```
...and finally, it is possible to combine several matchers in a single route:
```go
r.HandleFunc("/products", ProductsHandler).
Host("www.example.com").
Methods("GET").
Schemes("http")
```
Setting the same matching conditions again and again can be boring, so we have a way to group several routes that share the same requirements. We call it "subrouting".
For example, let's say we have several URLs that should only match when the host is `www.example.com`. Create a route for that host and get a "subrouter" from it:
```go
r := mux.NewRouter()
s := r.Host("www.example.com").Subrouter()
```
Then register routes in the subrouter:
```go
s.HandleFunc("/products/", ProductsHandler)
s.HandleFunc("/products/{key}", ProductHandler)
s.HandleFunc("/articles/{category}/{id:[0-9]+}"), ArticleHandler)
```
The three URL paths we registered above will only be tested if the domain is `www.example.com`, because the subrouter is tested first. This is not only convenient, but also optimizes request matching. You can create subrouters combining any attribute matchers accepted by a route.
Subrouters can be used to create domain or path "namespaces": you define subrouters in a central place and then parts of the app can register its paths relatively to a given subrouter.
There's one more thing about subroutes. When a subrouter has a path prefix, the inner routes use it as base for their paths:
```go
r := mux.NewRouter()
s := r.PathPrefix("/products").Subrouter()
// "/products/"
s.HandleFunc("/", ProductsHandler)
// "/products/{key}/"
s.HandleFunc("/{key}/", ProductHandler)
// "/products/{key}/details"
s.HandleFunc("/{key}/details", ProductDetailsHandler)
```
Now let's see how to build registered URLs.
Routes can be named. All routes that define a name can have their URLs built, or "reversed". We define a name calling `Name()` on a route. For example:
```go
r := mux.NewRouter()
r.HandleFunc("/articles/{category}/{id:[0-9]+}", ArticleHandler).
Name("article")
```
To build a URL, get the route and call the `URL()` method, passing a sequence of key/value pairs for the route variables. For the previous route, we would do:
```go
url, err := r.Get("article").URL("category", "technology", "id", "42")
```
...and the result will be a `url.URL` with the following path:
```
"/articles/technology/42"
```
This also works for host variables:
```go
r := mux.NewRouter()
r.Host("{subdomain}.domain.com").
Path("/articles/{category}/{id:[0-9]+}").
HandlerFunc(ArticleHandler).
Name("article")
// url.String() will be "http://news.domain.com/articles/technology/42"
url, err := r.Get("article").URL("subdomain", "news",
"category", "technology",
"id", "42")
```
All variables defined in the route are required, and their values must conform to the corresponding patterns. These requirements guarantee that a generated URL will always match a registered route -- the only exception is for explicitly defined "build-only" routes which never match.
Regex support also exists for matching Headers within a route. For example, we could do:
```go
r.HeadersRegexp("Content-Type", "application/(text|json)")
```
...and the route will match both requests with a Content-Type of `application/json` as well as `application/text`
There's also a way to build only the URL host or path for a route: use the methods `URLHost()` or `URLPath()` instead. For the previous route, we would do:
```go
// "http://news.domain.com/"
host, err := r.Get("article").URLHost("subdomain", "news")
// "/articles/technology/42"
path, err := r.Get("article").URLPath("category", "technology", "id", "42")
```
And if you use subrouters, host and path defined separately can be built as well:
```go
r := mux.NewRouter()
s := r.Host("{subdomain}.domain.com").Subrouter()
s.Path("/articles/{category}/{id:[0-9]+}").
HandlerFunc(ArticleHandler).
Name("article")
// "http://news.domain.com/articles/technology/42"
url, err := r.Get("article").URL("subdomain", "news",
"category", "technology",
"id", "42")
```
## Full Example
Here's a complete, runnable example of a small `mux` based server:
```go
package main
import (
"net/http"
"github.com/gorilla/mux"
)
func YourHandler(w http.ResponseWriter, r *http.Request) {
w.Write([]byte("Gorilla!\n"))
}
func main() {
r := mux.NewRouter()
// Routes consist of a path and a handler function.
r.HandleFunc("/", YourHandler)
// Bind to a port and pass our router in
http.ListenAndServe(":8000", r)
}
```
## License
BSD licensed. See the LICENSE file for details.
Generated Vendored Executable → Regular
View File
Generated Vendored Executable → Regular
+13 -6
View File
@@ -60,8 +60,8 @@ Routes can also be restricted to a domain or subdomain. Just define a host
pattern to be matched. They can also have variables: pattern to be matched. They can also have variables:
r := mux.NewRouter() r := mux.NewRouter()
// Only matches if domain is "www.domain.com". // Only matches if domain is "www.example.com".
r.Host("www.domain.com") r.Host("www.example.com")
// Matches a dynamic subdomain. // Matches a dynamic subdomain.
r.Host("{subdomain:[a-z]+}.domain.com") r.Host("{subdomain:[a-z]+}.domain.com")
@@ -94,7 +94,7 @@ There are several other matchers that can be added. To match path prefixes:
...and finally, it is possible to combine several matchers in a single route: ...and finally, it is possible to combine several matchers in a single route:
r.HandleFunc("/products", ProductsHandler). r.HandleFunc("/products", ProductsHandler).
Host("www.domain.com"). Host("www.example.com").
Methods("GET"). Methods("GET").
Schemes("http") Schemes("http")
@@ -103,11 +103,11 @@ a way to group several routes that share the same requirements.
We call it "subrouting". We call it "subrouting".
For example, let's say we have several URLs that should only match when the For example, let's say we have several URLs that should only match when the
host is "www.domain.com". Create a route for that host and get a "subrouter" host is "www.example.com". Create a route for that host and get a "subrouter"
from it: from it:
r := mux.NewRouter() r := mux.NewRouter()
s := r.Host("www.domain.com").Subrouter() s := r.Host("www.example.com").Subrouter()
Then register routes in the subrouter: Then register routes in the subrouter:
@@ -116,7 +116,7 @@ Then register routes in the subrouter:
s.HandleFunc("/articles/{category}/{id:[0-9]+}"), ArticleHandler) s.HandleFunc("/articles/{category}/{id:[0-9]+}"), ArticleHandler)
The three URL paths we registered above will only be tested if the domain is The three URL paths we registered above will only be tested if the domain is
"www.domain.com", because the subrouter is tested first. This is not "www.example.com", because the subrouter is tested first. This is not
only convenient, but also optimizes request matching. You can create only convenient, but also optimizes request matching. You can create
subrouters combining any attribute matchers accepted by a route. subrouters combining any attribute matchers accepted by a route.
@@ -172,6 +172,13 @@ conform to the corresponding patterns. These requirements guarantee that a
generated URL will always match a registered route -- the only exception is generated URL will always match a registered route -- the only exception is
for explicitly defined "build-only" routes which never match. for explicitly defined "build-only" routes which never match.
Regex support also exists for matching Headers within a route. For example, we could do:
r.HeadersRegexp("Content-Type", "application/(text|json)")
...and the route will match both requests with a Content-Type of `application/json` as well as
`application/text`
There's also a way to build only the URL host or path for a route: There's also a way to build only the URL host or path for a route:
use the methods URLHost() or URLPath() instead. For the previous route, use the methods URLHost() or URLPath() instead. For the previous route,
we would do: we would do:
Generated Vendored Executable → Regular
+128 -13
View File
@@ -5,9 +5,11 @@
package mux package mux
import ( import (
"errors"
"fmt" "fmt"
"net/http" "net/http"
"path" "path"
"regexp"
"github.com/gorilla/context" "github.com/gorilla/context"
) )
@@ -57,6 +59,12 @@ func (r *Router) Match(req *http.Request, match *RouteMatch) bool {
return true return true
} }
} }
// Closest match for a router (includes sub-routers)
if r.NotFoundHandler != nil {
match.Handler = r.NotFoundHandler
return true
}
return false return false
} }
@@ -68,7 +76,7 @@ func (r *Router) ServeHTTP(w http.ResponseWriter, req *http.Request) {
// Clean path to canonical form and redirect. // Clean path to canonical form and redirect.
if p := cleanPath(req.URL.Path); p != req.URL.Path { if p := cleanPath(req.URL.Path); p != req.URL.Path {
// Added 3 lines (Philip Schlump) - It was droping the query string and #whatever from query. // Added 3 lines (Philip Schlump) - It was dropping the query string and #whatever from query.
// This matches with fix in go 1.2 r.c. 4 for same problem. Go Issue: // This matches with fix in go 1.2 r.c. 4 for same problem. Go Issue:
// http://code.google.com/p/go/issues/detail?id=5252 // http://code.google.com/p/go/issues/detail?id=5252
url := *req.URL url := *req.URL
@@ -87,10 +95,7 @@ func (r *Router) ServeHTTP(w http.ResponseWriter, req *http.Request) {
setCurrentRoute(req, match.Route) setCurrentRoute(req, match.Route)
} }
if handler == nil { if handler == nil {
handler = r.NotFoundHandler handler = http.NotFoundHandler()
if handler == nil {
handler = http.NotFoundHandler()
}
} }
if !r.KeepContext { if !r.KeepContext {
defer context.Clear(req) defer context.Clear(req)
@@ -237,6 +242,52 @@ func (r *Router) BuildVarsFunc(f BuildVarsFunc) *Route {
return r.NewRoute().BuildVarsFunc(f) return r.NewRoute().BuildVarsFunc(f)
} }
// Walk walks the router and all its sub-routers, calling walkFn for each route
// in the tree. The routes are walked in the order they were added. Sub-routers
// are explored depth-first.
func (r *Router) Walk(walkFn WalkFunc) error {
return r.walk(walkFn, []*Route{})
}
// SkipRouter is used as a return value from WalkFuncs to indicate that the
// router that walk is about to descend down to should be skipped.
var SkipRouter = errors.New("skip this router")
// WalkFunc is the type of the function called for each route visited by Walk.
// At every invocation, it is given the current route, and the current router,
// and a list of ancestor routes that lead to the current route.
type WalkFunc func(route *Route, router *Router, ancestors []*Route) error
func (r *Router) walk(walkFn WalkFunc, ancestors []*Route) error {
for _, t := range r.routes {
if t.regexp == nil || t.regexp.path == nil || t.regexp.path.template == "" {
continue
}
err := walkFn(t, r, ancestors)
if err == SkipRouter {
continue
}
for _, sr := range t.matchers {
if h, ok := sr.(*Router); ok {
err := h.walk(walkFn, ancestors)
if err != nil {
return err
}
}
}
if h, ok := t.handler.(*Router); ok {
ancestors = append(ancestors, t)
err := h.walk(walkFn, ancestors)
if err != nil {
return err
}
ancestors = ancestors[:len(ancestors)-1]
}
}
return nil
}
// ---------------------------------------------------------------------------- // ----------------------------------------------------------------------------
// Context // Context
// ---------------------------------------------------------------------------- // ----------------------------------------------------------------------------
@@ -264,6 +315,10 @@ func Vars(r *http.Request) map[string]string {
} }
// CurrentRoute returns the matched route for the current request, if any. // CurrentRoute returns the matched route for the current request, if any.
// This only works when called inside the handler of the matched route
// because the matched route is stored in the request context which is cleared
// after the handler returns, unless the KeepContext option is set on the
// Router.
func CurrentRoute(r *http.Request) *Route { func CurrentRoute(r *http.Request) *Route {
if rv := context.Get(r, routeKey); rv != nil { if rv := context.Get(r, routeKey); rv != nil {
return rv.(*Route) return rv.(*Route)
@@ -272,11 +327,15 @@ func CurrentRoute(r *http.Request) *Route {
} }
func setVars(r *http.Request, val interface{}) { func setVars(r *http.Request, val interface{}) {
context.Set(r, varsKey, val) if val != nil {
context.Set(r, varsKey, val)
}
} }
func setCurrentRoute(r *http.Request, val interface{}) { func setCurrentRoute(r *http.Request, val interface{}) {
context.Set(r, routeKey, val) if val != nil {
context.Set(r, routeKey, val)
}
} }
// ---------------------------------------------------------------------------- // ----------------------------------------------------------------------------
@@ -313,13 +372,24 @@ func uniqueVars(s1, s2 []string) error {
return nil return nil
} }
// mapFromPairs converts variadic string parameters to a string map. // checkPairs returns the count of strings passed in, and an error if
func mapFromPairs(pairs ...string) (map[string]string, error) { // the count is not an even number.
func checkPairs(pairs ...string) (int, error) {
length := len(pairs) length := len(pairs)
if length%2 != 0 { if length%2 != 0 {
return nil, fmt.Errorf( return length, fmt.Errorf(
"mux: number of parameters must be multiple of 2, got %v", pairs) "mux: number of parameters must be multiple of 2, got %v", pairs)
} }
return length, nil
}
// mapFromPairsToString converts variadic string parameters to a
// string to string map.
func mapFromPairsToString(pairs ...string) (map[string]string, error) {
length, err := checkPairs(pairs...)
if err != nil {
return nil, err
}
m := make(map[string]string, length/2) m := make(map[string]string, length/2)
for i := 0; i < length; i += 2 { for i := 0; i < length; i += 2 {
m[pairs[i]] = pairs[i+1] m[pairs[i]] = pairs[i+1]
@@ -327,6 +397,24 @@ func mapFromPairs(pairs ...string) (map[string]string, error) {
return m, nil return m, nil
} }
// mapFromPairsToRegex converts variadic string paramers to a
// string to regex map.
func mapFromPairsToRegex(pairs ...string) (map[string]*regexp.Regexp, error) {
length, err := checkPairs(pairs...)
if err != nil {
return nil, err
}
m := make(map[string]*regexp.Regexp, length/2)
for i := 0; i < length; i += 2 {
regex, err := regexp.Compile(pairs[i+1])
if err != nil {
return nil, err
}
m[pairs[i]] = regex
}
return m, nil
}
// matchInArray returns true if the given string value is in the array. // matchInArray returns true if the given string value is in the array.
func matchInArray(arr []string, value string) bool { func matchInArray(arr []string, value string) bool {
for _, v := range arr { for _, v := range arr {
@@ -337,9 +425,8 @@ func matchInArray(arr []string, value string) bool {
return false return false
} }
// matchMap returns true if the given key/value pairs exist in a given map. // matchMapWithString returns true if the given key/value pairs exist in a given map.
func matchMap(toCheck map[string]string, toMatch map[string][]string, func matchMapWithString(toCheck map[string]string, toMatch map[string][]string, canonicalKey bool) bool {
canonicalKey bool) bool {
for k, v := range toCheck { for k, v := range toCheck {
// Check if key exists. // Check if key exists.
if canonicalKey { if canonicalKey {
@@ -364,3 +451,31 @@ func matchMap(toCheck map[string]string, toMatch map[string][]string,
} }
return true return true
} }
// matchMapWithRegex returns true if the given key/value pairs exist in a given map compiled against
// the given regex
func matchMapWithRegex(toCheck map[string]*regexp.Regexp, toMatch map[string][]string, canonicalKey bool) bool {
for k, v := range toCheck {
// Check if key exists.
if canonicalKey {
k = http.CanonicalHeaderKey(k)
}
if values := toMatch[k]; values == nil {
return false
} else if v != nil {
// If value was defined as an empty string we only check that the
// key exists. Otherwise we also check for equality.
valueExists := false
for _, value := range values {
if v.MatchString(value) {
valueExists = true
break
}
}
if !valueExists {
return false
}
}
}
return true
}
Generated Vendored Executable → Regular
+346
View File
@@ -7,11 +7,24 @@ package mux
import ( import (
"fmt" "fmt"
"net/http" "net/http"
"strings"
"testing" "testing"
"github.com/gorilla/context" "github.com/gorilla/context"
) )
func (r *Route) GoString() string {
matchers := make([]string, len(r.matchers))
for i, m := range r.matchers {
matchers[i] = fmt.Sprintf("%#v", m)
}
return fmt.Sprintf("&Route{matchers:[]matcher{%s}}", strings.Join(matchers, ", "))
}
func (r *routeRegexp) GoString() string {
return fmt.Sprintf("&routeRegexp{template: %q, matchHost: %t, matchQuery: %t, strictSlash: %t, regexp: regexp.MustCompile(%q), reverse: %q, varsN: %v, varsR: %v", r.template, r.matchHost, r.matchQuery, r.strictSlash, r.regexp.String(), r.reverse, r.varsN, r.varsR)
}
type routeTest struct { type routeTest struct {
title string // title of the test title string // title of the test
route *Route // the route being tested route *Route // the route being tested
@@ -108,6 +121,15 @@ func TestHost(t *testing.T) {
path: "", path: "",
shouldMatch: true, shouldMatch: true,
}, },
{
title: "Host route with pattern, additional capturing group, match",
route: new(Route).Host("aaa.{v1:[a-z]{2}(b|c)}.ccc"),
request: newRequest("GET", "http://aaa.bbb.ccc/111/222/333"),
vars: map[string]string{"v1": "bbb"},
host: "aaa.bbb.ccc",
path: "",
shouldMatch: true,
},
{ {
title: "Host route with pattern, wrong host in request URL", title: "Host route with pattern, wrong host in request URL",
route: new(Route).Host("aaa.{v1:[a-z]{3}}.ccc"), route: new(Route).Host("aaa.{v1:[a-z]{3}}.ccc"),
@@ -135,6 +157,33 @@ func TestHost(t *testing.T) {
path: "", path: "",
shouldMatch: false, shouldMatch: false,
}, },
{
title: "Host route with hyphenated name and pattern, match",
route: new(Route).Host("aaa.{v-1:[a-z]{3}}.ccc"),
request: newRequest("GET", "http://aaa.bbb.ccc/111/222/333"),
vars: map[string]string{"v-1": "bbb"},
host: "aaa.bbb.ccc",
path: "",
shouldMatch: true,
},
{
title: "Host route with hyphenated name and pattern, additional capturing group, match",
route: new(Route).Host("aaa.{v-1:[a-z]{2}(b|c)}.ccc"),
request: newRequest("GET", "http://aaa.bbb.ccc/111/222/333"),
vars: map[string]string{"v-1": "bbb"},
host: "aaa.bbb.ccc",
path: "",
shouldMatch: true,
},
{
title: "Host route with multiple hyphenated names and patterns, match",
route: new(Route).Host("{v-1:[a-z]{3}}.{v-2:[a-z]{3}}.{v-3:[a-z]{3}}"),
request: newRequest("GET", "http://aaa.bbb.ccc/111/222/333"),
vars: map[string]string{"v-1": "aaa", "v-2": "bbb", "v-3": "ccc"},
host: "aaa.bbb.ccc",
path: "",
shouldMatch: true,
},
{ {
title: "Path route with single pattern with pipe, match", title: "Path route with single pattern with pipe, match",
route: new(Route).Path("/{category:a|b/c}"), route: new(Route).Path("/{category:a|b/c}"),
@@ -260,6 +309,42 @@ func TestPath(t *testing.T) {
path: "/111/222/333", path: "/111/222/333",
shouldMatch: false, shouldMatch: false,
}, },
{
title: "Path route with multiple patterns with pipe, match",
route: new(Route).Path("/{category:a|(b/c)}/{product}/{id:[0-9]+}"),
request: newRequest("GET", "http://localhost/a/product_name/1"),
vars: map[string]string{"category": "a", "product": "product_name", "id": "1"},
host: "",
path: "/a/product_name/1",
shouldMatch: true,
},
{
title: "Path route with hyphenated name and pattern, match",
route: new(Route).Path("/111/{v-1:[0-9]{3}}/333"),
request: newRequest("GET", "http://localhost/111/222/333"),
vars: map[string]string{"v-1": "222"},
host: "",
path: "/111/222/333",
shouldMatch: true,
},
{
title: "Path route with multiple hyphenated names and patterns, match",
route: new(Route).Path("/{v-1:[0-9]{3}}/{v-2:[0-9]{3}}/{v-3:[0-9]{3}}"),
request: newRequest("GET", "http://localhost/111/222/333"),
vars: map[string]string{"v-1": "111", "v-2": "222", "v-3": "333"},
host: "",
path: "/111/222/333",
shouldMatch: true,
},
{
title: "Path route with multiple hyphenated names and patterns with pipe, match",
route: new(Route).Path("/{product-category:a|(b/c)}/{product-name}/{product-id:[0-9]+}"),
request: newRequest("GET", "http://localhost/a/product_name/1"),
vars: map[string]string{"product-category": "a", "product-name": "product_name", "product-id": "1"},
host: "",
path: "/a/product_name/1",
shouldMatch: true,
},
} }
for _, test := range tests { for _, test := range tests {
@@ -434,6 +519,24 @@ func TestHeaders(t *testing.T) {
path: "", path: "",
shouldMatch: false, shouldMatch: false,
}, },
{
title: "Headers route, regex header values to match",
route: new(Route).Headers("foo", "ba[zr]"),
request: newRequestHeaders("GET", "http://localhost", map[string]string{"foo": "bar"}),
vars: map[string]string{},
host: "",
path: "",
shouldMatch: false,
},
{
title: "Headers route, regex header values to match",
route: new(Route).HeadersRegexp("foo", "ba[zr]"),
request: newRequestHeaders("GET", "http://localhost", map[string]string{"foo": "baz"}),
vars: map[string]string{},
host: "",
path: "",
shouldMatch: true,
},
} }
for _, test := range tests { for _, test := range tests {
@@ -552,6 +655,150 @@ func TestQueries(t *testing.T) {
path: "", path: "",
shouldMatch: false, shouldMatch: false,
}, },
{
title: "Queries route with regexp pattern with quantifier, match",
route: new(Route).Queries("foo", "{v1:[0-9]{1}}"),
request: newRequest("GET", "http://localhost?foo=1"),
vars: map[string]string{"v1": "1"},
host: "",
path: "",
shouldMatch: true,
},
{
title: "Queries route with regexp pattern with quantifier, additional variable in query string, match",
route: new(Route).Queries("foo", "{v1:[0-9]{1}}"),
request: newRequest("GET", "http://localhost?bar=2&foo=1"),
vars: map[string]string{"v1": "1"},
host: "",
path: "",
shouldMatch: true,
},
{
title: "Queries route with regexp pattern with quantifier, regexp does not match",
route: new(Route).Queries("foo", "{v1:[0-9]{1}}"),
request: newRequest("GET", "http://localhost?foo=12"),
vars: map[string]string{},
host: "",
path: "",
shouldMatch: false,
},
{
title: "Queries route with regexp pattern with quantifier, additional capturing group",
route: new(Route).Queries("foo", "{v1:[0-9]{1}(a|b)}"),
request: newRequest("GET", "http://localhost?foo=1a"),
vars: map[string]string{"v1": "1a"},
host: "",
path: "",
shouldMatch: true,
},
{
title: "Queries route with regexp pattern with quantifier, additional variable in query string, regexp does not match",
route: new(Route).Queries("foo", "{v1:[0-9]{1}}"),
request: newRequest("GET", "http://localhost?foo=12"),
vars: map[string]string{},
host: "",
path: "",
shouldMatch: false,
},
{
title: "Queries route with hyphenated name, match",
route: new(Route).Queries("foo", "{v-1}"),
request: newRequest("GET", "http://localhost?foo=bar"),
vars: map[string]string{"v-1": "bar"},
host: "",
path: "",
shouldMatch: true,
},
{
title: "Queries route with multiple hyphenated names, match",
route: new(Route).Queries("foo", "{v-1}", "baz", "{v-2}"),
request: newRequest("GET", "http://localhost?foo=bar&baz=ding"),
vars: map[string]string{"v-1": "bar", "v-2": "ding"},
host: "",
path: "",
shouldMatch: true,
},
{
title: "Queries route with hyphenate name and pattern, match",
route: new(Route).Queries("foo", "{v-1:[0-9]+}"),
request: newRequest("GET", "http://localhost?foo=10"),
vars: map[string]string{"v-1": "10"},
host: "",
path: "",
shouldMatch: true,
},
{
title: "Queries route with hyphenated name and pattern with quantifier, additional capturing group",
route: new(Route).Queries("foo", "{v-1:[0-9]{1}(a|b)}"),
request: newRequest("GET", "http://localhost?foo=1a"),
vars: map[string]string{"v-1": "1a"},
host: "",
path: "",
shouldMatch: true,
},
{
title: "Queries route with empty value, should match",
route: new(Route).Queries("foo", ""),
request: newRequest("GET", "http://localhost?foo=bar"),
vars: map[string]string{},
host: "",
path: "",
shouldMatch: true,
},
{
title: "Queries route with empty value and no parameter in request, should not match",
route: new(Route).Queries("foo", ""),
request: newRequest("GET", "http://localhost"),
vars: map[string]string{},
host: "",
path: "",
shouldMatch: false,
},
{
title: "Queries route with empty value and empty parameter in request, should match",
route: new(Route).Queries("foo", ""),
request: newRequest("GET", "http://localhost?foo="),
vars: map[string]string{},
host: "",
path: "",
shouldMatch: true,
},
{
title: "Queries route with overlapping value, should not match",
route: new(Route).Queries("foo", "bar"),
request: newRequest("GET", "http://localhost?foo=barfoo"),
vars: map[string]string{},
host: "",
path: "",
shouldMatch: false,
},
{
title: "Queries route with no parameter in request, should not match",
route: new(Route).Queries("foo", "{bar}"),
request: newRequest("GET", "http://localhost"),
vars: map[string]string{},
host: "",
path: "",
shouldMatch: false,
},
{
title: "Queries route with empty parameter in request, should match",
route: new(Route).Queries("foo", "{bar}"),
request: newRequest("GET", "http://localhost?foo="),
vars: map[string]string{"foo": ""},
host: "",
path: "",
shouldMatch: true,
},
{
title: "Queries route, bad submatch",
route: new(Route).Queries("foo", "bar", "baz", "ding"),
request: newRequest("GET", "http://localhost?fffoo=bar&baz=dingggg"),
vars: map[string]string{},
host: "",
path: "",
shouldMatch: false,
},
} }
for _, test := range tests { for _, test := range tests {
@@ -801,6 +1048,105 @@ func TestStrictSlash(t *testing.T) {
} }
} }
func TestWalkSingleDepth(t *testing.T) {
r0 := NewRouter()
r1 := NewRouter()
r2 := NewRouter()
r0.Path("/g")
r0.Path("/o")
r0.Path("/d").Handler(r1)
r0.Path("/r").Handler(r2)
r0.Path("/a")
r1.Path("/z")
r1.Path("/i")
r1.Path("/l")
r1.Path("/l")
r2.Path("/i")
r2.Path("/l")
r2.Path("/l")
paths := []string{"g", "o", "r", "i", "l", "l", "a"}
depths := []int{0, 0, 0, 1, 1, 1, 0}
i := 0
err := r0.Walk(func(route *Route, router *Router, ancestors []*Route) error {
matcher := route.matchers[0].(*routeRegexp)
if matcher.template == "/d" {
return SkipRouter
}
if len(ancestors) != depths[i] {
t.Errorf(`Expected depth of %d at i = %d; got "%d"`, depths[i], i, len(ancestors))
}
if matcher.template != "/"+paths[i] {
t.Errorf(`Expected "/%s" at i = %d; got "%s"`, paths[i], i, matcher.template)
}
i++
return nil
})
if err != nil {
panic(err)
}
if i != len(paths) {
t.Errorf("Expected %d routes, found %d", len(paths), i)
}
}
func TestWalkNested(t *testing.T) {
router := NewRouter()
g := router.Path("/g").Subrouter()
o := g.PathPrefix("/o").Subrouter()
r := o.PathPrefix("/r").Subrouter()
i := r.PathPrefix("/i").Subrouter()
l1 := i.PathPrefix("/l").Subrouter()
l2 := l1.PathPrefix("/l").Subrouter()
l2.Path("/a")
paths := []string{"/g", "/g/o", "/g/o/r", "/g/o/r/i", "/g/o/r/i/l", "/g/o/r/i/l/l", "/g/o/r/i/l/l/a"}
idx := 0
err := router.Walk(func(route *Route, router *Router, ancestors []*Route) error {
path := paths[idx]
tpl := route.regexp.path.template
if tpl != path {
t.Errorf(`Expected %s got %s`, path, tpl)
}
idx++
return nil
})
if err != nil {
panic(err)
}
if idx != len(paths) {
t.Errorf("Expected %d routes, found %d", len(paths), idx)
}
}
func TestSubrouterErrorHandling(t *testing.T) {
superRouterCalled := false
subRouterCalled := false
router := NewRouter()
router.NotFoundHandler = http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) {
superRouterCalled = true
})
subRouter := router.PathPrefix("/bign8").Subrouter()
subRouter.NotFoundHandler = http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) {
subRouterCalled = true
})
req, _ := http.NewRequest("GET", "http://localhost/bign8/was/here", nil)
router.ServeHTTP(NewRecorder(), req)
if superRouterCalled {
t.Error("Super router 404 handler called when sub-router 404 handler is available.")
}
if !subRouterCalled {
t.Error("Sub-router 404 handler was not called.")
}
}
// ---------------------------------------------------------------------------- // ----------------------------------------------------------------------------
// Helpers // Helpers
// ---------------------------------------------------------------------------- // ----------------------------------------------------------------------------
Generated Vendored Executable → Regular
+3 -3
View File
@@ -545,7 +545,7 @@ func TestMatchedRouteName(t *testing.T) {
router := NewRouter() router := NewRouter()
route := router.NewRoute().Path("/products/").Name(routeName) route := router.NewRoute().Path("/products/").Name(routeName)
url := "http://www.domain.com/products/" url := "http://www.example.com/products/"
request, _ := http.NewRequest("GET", url, nil) request, _ := http.NewRequest("GET", url, nil)
var rv RouteMatch var rv RouteMatch
ok := router.Match(request, &rv) ok := router.Match(request, &rv)
@@ -563,10 +563,10 @@ func TestMatchedRouteName(t *testing.T) {
func TestSubRouting(t *testing.T) { func TestSubRouting(t *testing.T) {
// Example from docs. // Example from docs.
router := NewRouter() router := NewRouter()
subrouter := router.NewRoute().Host("www.domain.com").Subrouter() subrouter := router.NewRoute().Host("www.example.com").Subrouter()
route := subrouter.NewRoute().Path("/products/").Name("products") route := subrouter.NewRoute().Path("/products/").Name("products")
url := "http://www.domain.com/products/" url := "http://www.example.com/products/"
request, _ := http.NewRequest("GET", url, nil) request, _ := http.NewRequest("GET", url, nil)
var rv RouteMatch var rv RouteMatch
ok := router.Match(request, &rv) ok := router.Match(request, &rv)
Generated Vendored Executable → Regular
+62 -17
View File
@@ -10,6 +10,7 @@ import (
"net/http" "net/http"
"net/url" "net/url"
"regexp" "regexp"
"strconv"
"strings" "strings"
) )
@@ -34,8 +35,7 @@ func newRouteRegexp(tpl string, matchHost, matchPrefix, matchQuery, strictSlash
// Now let's parse it. // Now let's parse it.
defaultPattern := "[^/]+" defaultPattern := "[^/]+"
if matchQuery { if matchQuery {
defaultPattern = "[^?&]+" defaultPattern = "[^?&]*"
matchPrefix = true
} else if matchHost { } else if matchHost {
defaultPattern = "[^.]+" defaultPattern = "[^.]+"
matchPrefix = false matchPrefix = false
@@ -53,9 +53,7 @@ func newRouteRegexp(tpl string, matchHost, matchPrefix, matchQuery, strictSlash
varsN := make([]string, len(idxs)/2) varsN := make([]string, len(idxs)/2)
varsR := make([]*regexp.Regexp, len(idxs)/2) varsR := make([]*regexp.Regexp, len(idxs)/2)
pattern := bytes.NewBufferString("") pattern := bytes.NewBufferString("")
if !matchQuery { pattern.WriteByte('^')
pattern.WriteByte('^')
}
reverse := bytes.NewBufferString("") reverse := bytes.NewBufferString("")
var end int var end int
var err error var err error
@@ -75,12 +73,14 @@ func newRouteRegexp(tpl string, matchHost, matchPrefix, matchQuery, strictSlash
tpl[idxs[i]:end]) tpl[idxs[i]:end])
} }
// Build the regexp pattern. // Build the regexp pattern.
fmt.Fprintf(pattern, "%s(%s)", regexp.QuoteMeta(raw), patt) varIdx := i / 2
fmt.Fprintf(pattern, "%s(?P<%s>%s)", regexp.QuoteMeta(raw), varGroupName(varIdx), patt)
// Build the reverse template. // Build the reverse template.
fmt.Fprintf(reverse, "%s%%s", raw) fmt.Fprintf(reverse, "%s%%s", raw)
// Append variable name and compiled pattern. // Append variable name and compiled pattern.
varsN[i/2] = name varsN[varIdx] = name
varsR[i/2], err = regexp.Compile(fmt.Sprintf("^%s$", patt)) varsR[varIdx], err = regexp.Compile(fmt.Sprintf("^%s$", patt))
if err != nil { if err != nil {
return nil, err return nil, err
} }
@@ -91,6 +91,12 @@ func newRouteRegexp(tpl string, matchHost, matchPrefix, matchQuery, strictSlash
if strictSlash { if strictSlash {
pattern.WriteString("[/]?") pattern.WriteString("[/]?")
} }
if matchQuery {
// Add the default pattern if the query value is empty
if queryVal := strings.SplitN(template, "=", 2)[1]; queryVal == "" {
pattern.WriteString(defaultPattern)
}
}
if !matchPrefix { if !matchPrefix {
pattern.WriteByte('$') pattern.WriteByte('$')
} }
@@ -141,7 +147,7 @@ type routeRegexp struct {
func (r *routeRegexp) Match(req *http.Request, match *RouteMatch) bool { func (r *routeRegexp) Match(req *http.Request, match *RouteMatch) bool {
if !r.matchHost { if !r.matchHost {
if r.matchQuery { if r.matchQuery {
return r.regexp.MatchString(req.URL.RawQuery) return r.matchQueryString(req)
} else { } else {
return r.regexp.MatchString(req.URL.Path) return r.regexp.MatchString(req.URL.Path)
} }
@@ -175,6 +181,26 @@ func (r *routeRegexp) url(values map[string]string) (string, error) {
return rv, nil return rv, nil
} }
// getUrlQuery returns a single query parameter from a request URL.
// For a URL with foo=bar&baz=ding, we return only the relevant key
// value pair for the routeRegexp.
func (r *routeRegexp) getUrlQuery(req *http.Request) string {
if !r.matchQuery {
return ""
}
templateKey := strings.SplitN(r.template, "=", 2)[0]
for key, vals := range req.URL.Query() {
if key == templateKey && len(vals) > 0 {
return key + "=" + vals[0]
}
}
return ""
}
func (r *routeRegexp) matchQueryString(req *http.Request) bool {
return r.regexp.MatchString(r.getUrlQuery(req))
}
// braceIndices returns the first level curly brace indices from a string. // braceIndices returns the first level curly brace indices from a string.
// It returns an error in case of unbalanced braces. // It returns an error in case of unbalanced braces.
func braceIndices(s string) ([]int, error) { func braceIndices(s string) ([]int, error) {
@@ -200,6 +226,11 @@ func braceIndices(s string) ([]int, error) {
return idxs, nil return idxs, nil
} }
// varGroupName builds a capturing group name for the indexed variable.
func varGroupName(idx int) string {
return "v" + strconv.Itoa(idx)
}
// ---------------------------------------------------------------------------- // ----------------------------------------------------------------------------
// routeRegexpGroup // routeRegexpGroup
// ---------------------------------------------------------------------------- // ----------------------------------------------------------------------------
@@ -217,8 +248,13 @@ func (v *routeRegexpGroup) setMatch(req *http.Request, m *RouteMatch, r *Route)
if v.host != nil { if v.host != nil {
hostVars := v.host.regexp.FindStringSubmatch(getHost(req)) hostVars := v.host.regexp.FindStringSubmatch(getHost(req))
if hostVars != nil { if hostVars != nil {
for k, v := range v.host.varsN { subexpNames := v.host.regexp.SubexpNames()
m.Vars[v] = hostVars[k+1] varName := 0
for i, name := range subexpNames[1:] {
if name != "" && name == varGroupName(varName) {
m.Vars[v.host.varsN[varName]] = hostVars[i+1]
varName++
}
} }
} }
} }
@@ -226,8 +262,13 @@ func (v *routeRegexpGroup) setMatch(req *http.Request, m *RouteMatch, r *Route)
if v.path != nil { if v.path != nil {
pathVars := v.path.regexp.FindStringSubmatch(req.URL.Path) pathVars := v.path.regexp.FindStringSubmatch(req.URL.Path)
if pathVars != nil { if pathVars != nil {
for k, v := range v.path.varsN { subexpNames := v.path.regexp.SubexpNames()
m.Vars[v] = pathVars[k+1] varName := 0
for i, name := range subexpNames[1:] {
if name != "" && name == varGroupName(varName) {
m.Vars[v.path.varsN[varName]] = pathVars[i+1]
varName++
}
} }
// Check if we should redirect. // Check if we should redirect.
if v.path.strictSlash { if v.path.strictSlash {
@@ -246,12 +287,16 @@ func (v *routeRegexpGroup) setMatch(req *http.Request, m *RouteMatch, r *Route)
} }
} }
// Store query string variables. // Store query string variables.
rawQuery := req.URL.RawQuery
for _, q := range v.queries { for _, q := range v.queries {
queryVars := q.regexp.FindStringSubmatch(rawQuery) queryVars := q.regexp.FindStringSubmatch(q.getUrlQuery(req))
if queryVars != nil { if queryVars != nil {
for k, v := range q.varsN { subexpNames := q.regexp.SubexpNames()
m.Vars[v] = queryVars[k+1] varName := 0
for i, name := range subexpNames[1:] {
if name != "" && name == varGroupName(varName) {
m.Vars[q.varsN[varName]] = queryVars[i+1]
varName++
}
} }
} }
} }
Generated Vendored Executable → Regular
+32 -8
View File
@@ -9,6 +9,7 @@ import (
"fmt" "fmt"
"net/http" "net/http"
"net/url" "net/url"
"regexp"
"strings" "strings"
) )
@@ -188,7 +189,7 @@ func (r *Route) addRegexpMatcher(tpl string, matchHost, matchPrefix, matchQuery
type headerMatcher map[string]string type headerMatcher map[string]string
func (m headerMatcher) Match(r *http.Request, match *RouteMatch) bool { func (m headerMatcher) Match(r *http.Request, match *RouteMatch) bool {
return matchMap(m, r.Header, true) return matchMapWithString(m, r.Header, true)
} }
// Headers adds a matcher for request header values. // Headers adds a matcher for request header values.
@@ -199,17 +200,40 @@ func (m headerMatcher) Match(r *http.Request, match *RouteMatch) bool {
// "X-Requested-With", "XMLHttpRequest") // "X-Requested-With", "XMLHttpRequest")
// //
// The above route will only match if both request header values match. // The above route will only match if both request header values match.
// // If the value is an empty string, it will match any value if the key is set.
// It the value is an empty string, it will match any value if the key is set.
func (r *Route) Headers(pairs ...string) *Route { func (r *Route) Headers(pairs ...string) *Route {
if r.err == nil { if r.err == nil {
var headers map[string]string var headers map[string]string
headers, r.err = mapFromPairs(pairs...) headers, r.err = mapFromPairsToString(pairs...)
return r.addMatcher(headerMatcher(headers)) return r.addMatcher(headerMatcher(headers))
} }
return r return r
} }
// headerRegexMatcher matches the request against the route given a regex for the header
type headerRegexMatcher map[string]*regexp.Regexp
func (m headerRegexMatcher) Match(r *http.Request, match *RouteMatch) bool {
return matchMapWithRegex(m, r.Header, true)
}
// Regular expressions can be used with headers as well.
// It accepts a sequence of key/value pairs, where the value has regex support. For example
// r := mux.NewRouter()
// r.HeadersRegexp("Content-Type", "application/(text|json)",
// "X-Requested-With", "XMLHttpRequest")
//
// The above route will only match if both the request header matches both regular expressions.
// It the value is an empty string, it will match any value if the key is set.
func (r *Route) HeadersRegexp(pairs ...string) *Route {
if r.err == nil {
var headers map[string]*regexp.Regexp
headers, r.err = mapFromPairsToRegex(pairs...)
return r.addMatcher(headerRegexMatcher(headers))
}
return r
}
// Host ----------------------------------------------------------------------- // Host -----------------------------------------------------------------------
// Host adds a matcher for the URL host. // Host adds a matcher for the URL host.
@@ -223,7 +247,7 @@ func (r *Route) Headers(pairs ...string) *Route {
// For example: // For example:
// //
// r := mux.NewRouter() // r := mux.NewRouter()
// r.Host("www.domain.com") // r.Host("www.example.com")
// r.Host("{subdomain}.domain.com") // r.Host("{subdomain}.domain.com")
// r.Host("{subdomain:[a-z]+}.domain.com") // r.Host("{subdomain:[a-z]+}.domain.com")
// //
@@ -336,7 +360,7 @@ func (r *Route) Queries(pairs ...string) *Route {
return nil return nil
} }
for i := 0; i < length; i += 2 { for i := 0; i < length; i += 2 {
if r.err = r.addRegexpMatcher(pairs[i]+"="+pairs[i+1], false, true, true); r.err != nil { if r.err = r.addRegexpMatcher(pairs[i]+"="+pairs[i+1], false, false, true); r.err != nil {
return r return r
} }
} }
@@ -382,7 +406,7 @@ func (r *Route) BuildVarsFunc(f BuildVarsFunc) *Route {
// It will test the inner routes only if the parent route matched. For example: // It will test the inner routes only if the parent route matched. For example:
// //
// r := mux.NewRouter() // r := mux.NewRouter()
// s := r.Host("www.domain.com").Subrouter() // s := r.Host("www.example.com").Subrouter()
// s.HandleFunc("/products/", ProductsHandler) // s.HandleFunc("/products/", ProductsHandler)
// s.HandleFunc("/products/{key}", ProductHandler) // s.HandleFunc("/products/{key}", ProductHandler)
// s.HandleFunc("/articles/{category}/{id:[0-9]+}"), ArticleHandler) // s.HandleFunc("/articles/{category}/{id:[0-9]+}"), ArticleHandler)
@@ -511,7 +535,7 @@ func (r *Route) URLPath(pairs ...string) (*url.URL, error) {
// prepareVars converts the route variable pairs into a map. If the route has a // prepareVars converts the route variable pairs into a map. If the route has a
// BuildVarsFunc, it is invoked. // BuildVarsFunc, it is invoked.
func (r *Route) prepareVars(pairs ...string) (map[string]string, error) { func (r *Route) prepareVars(pairs ...string) (map[string]string, error) {
m, err := mapFromPairs(pairs...) m, err := mapFromPairsToString(pairs...)
if err != nil { if err != nil {
return nil, err return nil, err
} }
Generated Vendored Executable → Regular
View File
Generated Vendored Executable → Regular
View File
Generated Vendored Executable → Regular
+2
View File
@@ -7,6 +7,8 @@ Gorilla WebSocket is a [Go](http://golang.org/) implementation of the
* [API Reference](http://godoc.org/github.com/gorilla/websocket) * [API Reference](http://godoc.org/github.com/gorilla/websocket)
* [Chat example](https://github.com/gorilla/websocket/tree/master/examples/chat) * [Chat example](https://github.com/gorilla/websocket/tree/master/examples/chat)
* [Command example](https://github.com/gorilla/websocket/tree/master/examples/command)
* [Client and server example](https://github.com/gorilla/websocket/tree/master/examples/echo)
* [File watch example](https://github.com/gorilla/websocket/tree/master/examples/filewatch) * [File watch example](https://github.com/gorilla/websocket/tree/master/examples/filewatch)
### Status ### Status
Generated Vendored Executable → Regular
View File
Generated Vendored Executable → Regular
+176 -95
View File
@@ -5,8 +5,10 @@
package websocket package websocket
import ( import (
"bufio"
"bytes" "bytes"
"crypto/tls" "crypto/tls"
"encoding/base64"
"errors" "errors"
"io" "io"
"io/ioutil" "io/ioutil"
@@ -30,50 +32,17 @@ var ErrBadHandshake = errors.New("websocket: bad handshake")
// If the WebSocket handshake fails, ErrBadHandshake is returned along with a // If the WebSocket handshake fails, ErrBadHandshake is returned along with a
// non-nil *http.Response so that callers can handle redirects, authentication, // non-nil *http.Response so that callers can handle redirects, authentication,
// etc. // etc.
//
// Deprecated: Use Dialer instead.
func NewClient(netConn net.Conn, u *url.URL, requestHeader http.Header, readBufSize, writeBufSize int) (c *Conn, response *http.Response, err error) { func NewClient(netConn net.Conn, u *url.URL, requestHeader http.Header, readBufSize, writeBufSize int) (c *Conn, response *http.Response, err error) {
challengeKey, err := generateChallengeKey() d := Dialer{
if err != nil { ReadBufferSize: readBufSize,
return nil, nil, err WriteBufferSize: writeBufSize,
NetDial: func(net, addr string) (net.Conn, error) {
return netConn, nil
},
} }
acceptKey := computeAcceptKey(challengeKey) return d.Dial(u.String(), requestHeader)
c = newConn(netConn, false, readBufSize, writeBufSize)
p := c.writeBuf[:0]
p = append(p, "GET "...)
p = append(p, u.RequestURI()...)
p = append(p, " HTTP/1.1\r\nHost: "...)
p = append(p, u.Host...)
// "Upgrade" is capitalized for servers that do not use case insensitive
// comparisons on header tokens.
p = append(p, "\r\nUpgrade: websocket\r\nConnection: Upgrade\r\nSec-WebSocket-Version: 13\r\nSec-WebSocket-Key: "...)
p = append(p, challengeKey...)
p = append(p, "\r\n"...)
for k, vs := range requestHeader {
for _, v := range vs {
p = append(p, k...)
p = append(p, ": "...)
p = append(p, v...)
p = append(p, "\r\n"...)
}
}
p = append(p, "\r\n"...)
if _, err := netConn.Write(p); err != nil {
return nil, nil, err
}
resp, err := http.ReadResponse(c.br, &http.Request{Method: "GET", URL: u})
if err != nil {
return nil, nil, err
}
if resp.StatusCode != 101 ||
!strings.EqualFold(resp.Header.Get("Upgrade"), "websocket") ||
!strings.EqualFold(resp.Header.Get("Connection"), "upgrade") ||
resp.Header.Get("Sec-Websocket-Accept") != acceptKey {
return nil, resp, ErrBadHandshake
}
c.subprotocol = resp.Header.Get("Sec-Websocket-Protocol")
return c, resp, nil
} }
// A Dialer contains options for connecting to WebSocket server. // A Dialer contains options for connecting to WebSocket server.
@@ -82,6 +51,12 @@ type Dialer struct {
// NetDial is nil, net.Dial is used. // NetDial is nil, net.Dial is used.
NetDial func(network, addr string) (net.Conn, error) NetDial func(network, addr string) (net.Conn, error)
// Proxy specifies a function to return a proxy for a given
// Request. If the function returns a non-nil error, the
// request is aborted with the provided error.
// If Proxy is nil or returns a nil *URL, no proxy is used.
Proxy func(*http.Request) (*url.URL, error)
// TLSClientConfig specifies the TLS configuration to use with tls.Client. // TLSClientConfig specifies the TLS configuration to use with tls.Client.
// If nil, the default configuration is used. // If nil, the default configuration is used.
TLSClientConfig *tls.Config TLSClientConfig *tls.Config
@@ -99,17 +74,15 @@ type Dialer struct {
var errMalformedURL = errors.New("malformed ws or wss URL") var errMalformedURL = errors.New("malformed ws or wss URL")
// parseURL parses the URL. The url.Parse function is not used here because // parseURL parses the URL.
// url.Parse mangles the path. //
// This function is a replacement for the standard library url.Parse function.
// In Go 1.4 and earlier, url.Parse loses information from the path.
func parseURL(s string) (*url.URL, error) { func parseURL(s string) (*url.URL, error) {
// From the RFC: // From the RFC:
// //
// ws-URI = "ws:" "//" host [ ":" port ] path [ "?" query ] // ws-URI = "ws:" "//" host [ ":" port ] path [ "?" query ]
// wss-URI = "wss:" "//" host [ ":" port ] path [ "?" query ] // wss-URI = "wss:" "//" host [ ":" port ] path [ "?" query ]
//
// We don't use the net/url parser here because the dialer interface does
// not provide a way for applications to work around percent deocding in
// the net/url parser.
var u url.URL var u url.URL
switch { switch {
@@ -123,15 +96,23 @@ func parseURL(s string) (*url.URL, error) {
return nil, errMalformedURL return nil, errMalformedURL
} }
u.Host = s if i := strings.Index(s, "?"); i >= 0 {
u.Opaque = "/" u.RawQuery = s[i+1:]
if i := strings.Index(s, "/"); i >= 0 { s = s[:i]
u.Host = s[:i]
u.Opaque = s[i:]
} }
if i := strings.Index(s, "/"); i >= 0 {
u.Opaque = s[i:]
s = s[:i]
} else {
u.Opaque = "/"
}
u.Host = s
if strings.Contains(u.Host, "@") { if strings.Contains(u.Host, "@") {
// WebSocket URIs do not contain user information. // Don't bother parsing user information because user information is
// not allowed in websocket URIs.
return nil, errMalformedURL return nil, errMalformedURL
} }
@@ -144,9 +125,12 @@ func hostPortNoPort(u *url.URL) (hostPort, hostNoPort string) {
if i := strings.LastIndex(u.Host, ":"); i > strings.LastIndex(u.Host, "]") { if i := strings.LastIndex(u.Host, ":"); i > strings.LastIndex(u.Host, "]") {
hostNoPort = hostNoPort[:i] hostNoPort = hostNoPort[:i]
} else { } else {
if u.Scheme == "wss" { switch u.Scheme {
case "wss":
hostPort += ":443" hostPort += ":443"
} else { case "https":
hostPort += ":443"
default:
hostPort += ":80" hostPort += ":80"
} }
} }
@@ -154,7 +138,9 @@ func hostPortNoPort(u *url.URL) (hostPort, hostNoPort string) {
} }
// DefaultDialer is a dialer with all fields set to the default zero values. // DefaultDialer is a dialer with all fields set to the default zero values.
var DefaultDialer *Dialer var DefaultDialer = &Dialer{
Proxy: http.ProxyFromEnvironment,
}
// Dial creates a new client connection. Use requestHeader to specify the // Dial creates a new client connection. Use requestHeader to specify the
// origin (Origin), subprotocols (Sec-WebSocket-Protocol) and cookies (Cookie). // origin (Origin), subprotocols (Sec-WebSocket-Protocol) and cookies (Cookie).
@@ -166,15 +152,91 @@ var DefaultDialer *Dialer
// etcetera. The response body may not contain the entire response and does not // etcetera. The response body may not contain the entire response and does not
// need to be closed by the application. // need to be closed by the application.
func (d *Dialer) Dial(urlStr string, requestHeader http.Header) (*Conn, *http.Response, error) { func (d *Dialer) Dial(urlStr string, requestHeader http.Header) (*Conn, *http.Response, error) {
if d == nil {
d = &Dialer{
Proxy: http.ProxyFromEnvironment,
}
}
challengeKey, err := generateChallengeKey()
if err != nil {
return nil, nil, err
}
u, err := parseURL(urlStr) u, err := parseURL(urlStr)
if err != nil { if err != nil {
return nil, nil, err return nil, nil, err
} }
switch u.Scheme {
case "ws":
u.Scheme = "http"
case "wss":
u.Scheme = "https"
default:
return nil, nil, errMalformedURL
}
if u.User != nil {
// User name and password are not allowed in websocket URIs.
return nil, nil, errMalformedURL
}
req := &http.Request{
Method: "GET",
URL: u,
Proto: "HTTP/1.1",
ProtoMajor: 1,
ProtoMinor: 1,
Header: make(http.Header),
Host: u.Host,
}
// Set the request headers using the capitalization for names and values in
// RFC examples. Although the capitalization shouldn't matter, there are
// servers that depend on it. The Header.Set method is not used because the
// method canonicalizes the header names.
req.Header["Upgrade"] = []string{"websocket"}
req.Header["Connection"] = []string{"Upgrade"}
req.Header["Sec-WebSocket-Key"] = []string{challengeKey}
req.Header["Sec-WebSocket-Version"] = []string{"13"}
if len(d.Subprotocols) > 0 {
req.Header["Sec-WebSocket-Protocol"] = []string{strings.Join(d.Subprotocols, ", ")}
}
for k, vs := range requestHeader {
switch {
case k == "Host":
if len(vs) > 0 {
req.Host = vs[0]
}
case k == "Upgrade" ||
k == "Connection" ||
k == "Sec-Websocket-Key" ||
k == "Sec-Websocket-Version" ||
(k == "Sec-Websocket-Protocol" && len(d.Subprotocols) > 0):
return nil, nil, errors.New("websocket: duplicate header not allowed: " + k)
default:
req.Header[k] = vs
}
}
hostPort, hostNoPort := hostPortNoPort(u) hostPort, hostNoPort := hostPortNoPort(u)
if d == nil { var proxyURL *url.URL
d = &Dialer{} // Check wether the proxy method has been configured
if d.Proxy != nil {
proxyURL, err = d.Proxy(req)
}
if err != nil {
return nil, nil, err
}
var targetHostPort string
if proxyURL != nil {
targetHostPort, _ = hostPortNoPort(proxyURL)
} else {
targetHostPort = hostPort
} }
var deadline time.Time var deadline time.Time
@@ -188,7 +250,7 @@ func (d *Dialer) Dial(urlStr string, requestHeader http.Header) (*Conn, *http.Re
netDial = netDialer.Dial netDial = netDialer.Dial
} }
netConn, err := netDial("tcp", hostPort) netConn, err := netDial("tcp", targetHostPort)
if err != nil { if err != nil {
return nil, nil, err return nil, nil, err
} }
@@ -203,7 +265,39 @@ func (d *Dialer) Dial(urlStr string, requestHeader http.Header) (*Conn, *http.Re
return nil, nil, err return nil, nil, err
} }
if u.Scheme == "wss" { if proxyURL != nil {
connectHeader := make(http.Header)
if user := proxyURL.User; user != nil {
proxyUser := user.Username()
if proxyPassword, passwordSet := user.Password(); passwordSet {
credential := base64.StdEncoding.EncodeToString([]byte(proxyUser + ":" + proxyPassword))
connectHeader.Set("Proxy-Authorization", "Basic "+credential)
}
}
connectReq := &http.Request{
Method: "CONNECT",
URL: &url.URL{Opaque: hostPort},
Host: hostPort,
Header: connectHeader,
}
connectReq.Write(netConn)
// Read response.
// Okay to use and discard buffered reader here, because
// TLS server will not speak until spoken to.
br := bufio.NewReader(netConn)
resp, err := http.ReadResponse(br, connectReq)
if err != nil {
return nil, nil, err
}
if resp.StatusCode != 200 {
f := strings.SplitN(resp.Status, " ", 2)
return nil, nil, errors.New(f[1])
}
}
if u.Scheme == "https" {
cfg := d.TLSClientConfig cfg := d.TLSClientConfig
if cfg == nil { if cfg == nil {
cfg = &tls.Config{ServerName: hostNoPort} cfg = &tls.Config{ServerName: hostNoPort}
@@ -224,44 +318,31 @@ func (d *Dialer) Dial(urlStr string, requestHeader http.Header) (*Conn, *http.Re
} }
} }
if len(d.Subprotocols) > 0 { conn := newConn(netConn, false, d.ReadBufferSize, d.WriteBufferSize)
h := http.Header{}
for k, v := range requestHeader { if err := req.Write(netConn); err != nil {
h[k] = v return nil, nil, err
}
h.Set("Sec-Websocket-Protocol", strings.Join(d.Subprotocols, ", "))
requestHeader = h
} }
if len(requestHeader["Host"]) > 0 { resp, err := http.ReadResponse(conn.br, req)
// This can be used to supply a Host: header which is different from
// the dial address.
u.Host = requestHeader.Get("Host")
// Drop "Host" header
h := http.Header{}
for k, v := range requestHeader {
if k == "Host" {
continue
}
h[k] = v
}
requestHeader = h
}
conn, resp, err := NewClient(netConn, u, requestHeader, d.ReadBufferSize, d.WriteBufferSize)
if err != nil { if err != nil {
if err == ErrBadHandshake { return nil, nil, err
// Before closing the network connection on return from this
// function, slurp up some of the response to aid application
// debugging.
buf := make([]byte, 1024)
n, _ := io.ReadFull(resp.Body, buf)
resp.Body = ioutil.NopCloser(bytes.NewReader(buf[:n]))
}
return nil, resp, err
} }
if resp.StatusCode != 101 ||
!strings.EqualFold(resp.Header.Get("Upgrade"), "websocket") ||
!strings.EqualFold(resp.Header.Get("Connection"), "upgrade") ||
resp.Header.Get("Sec-Websocket-Accept") != computeAcceptKey(challengeKey) {
// Before closing the network connection on return from this
// function, slurp up some of the response to aid application
// debugging.
buf := make([]byte, 1024)
n, _ := io.ReadFull(resp.Body, buf)
resp.Body = ioutil.NopCloser(bytes.NewReader(buf[:n]))
return nil, resp, ErrBadHandshake
}
resp.Body = ioutil.NopCloser(bytes.NewReader([]byte{}))
conn.subprotocol = resp.Header.Get("Sec-Websocket-Protocol")
netConn.SetDeadline(time.Time{}) netConn.SetDeadline(time.Time{})
netConn = nil // to avoid close in defer. netConn = nil // to avoid close in defer.
Generated Vendored Executable → Regular
+139 -11
View File
@@ -7,6 +7,7 @@ package websocket
import ( import (
"crypto/tls" "crypto/tls"
"crypto/x509" "crypto/x509"
"encoding/base64"
"io" "io"
"io/ioutil" "io/ioutil"
"net" "net"
@@ -41,9 +42,16 @@ type cstServer struct {
URL string URL string
} }
const (
cstPath = "/a/b"
cstRawQuery = "x=y"
cstRequestURI = cstPath + "?" + cstRawQuery
)
func newServer(t *testing.T) *cstServer { func newServer(t *testing.T) *cstServer {
var s cstServer var s cstServer
s.Server = httptest.NewServer(cstHandler{t}) s.Server = httptest.NewServer(cstHandler{t})
s.Server.URL += cstRequestURI
s.URL = makeWsProto(s.Server.URL) s.URL = makeWsProto(s.Server.URL)
return &s return &s
} }
@@ -51,14 +59,20 @@ func newServer(t *testing.T) *cstServer {
func newTLSServer(t *testing.T) *cstServer { func newTLSServer(t *testing.T) *cstServer {
var s cstServer var s cstServer
s.Server = httptest.NewTLSServer(cstHandler{t}) s.Server = httptest.NewTLSServer(cstHandler{t})
s.Server.URL += cstRequestURI
s.URL = makeWsProto(s.Server.URL) s.URL = makeWsProto(s.Server.URL)
return &s return &s
} }
func (t cstHandler) ServeHTTP(w http.ResponseWriter, r *http.Request) { func (t cstHandler) ServeHTTP(w http.ResponseWriter, r *http.Request) {
if r.Method != "GET" { if r.URL.Path != cstPath {
t.Logf("method %s not allowed", r.Method) t.Logf("path=%v, want %v", r.URL.Path, cstPath)
http.Error(w, "method not allowed", 405) http.Error(w, "bad path", 400)
return
}
if r.URL.RawQuery != cstRawQuery {
t.Logf("query=%v, want %v", r.URL.RawQuery, cstRawQuery)
http.Error(w, "bad path", 400)
return return
} }
subprotos := Subprotocols(r) subprotos := Subprotocols(r)
@@ -123,6 +137,85 @@ func sendRecv(t *testing.T, ws *Conn) {
} }
} }
func TestProxyDial(t *testing.T) {
s := newServer(t)
defer s.Close()
surl, _ := url.Parse(s.URL)
cstDialer.Proxy = http.ProxyURL(surl)
connect := false
origHandler := s.Server.Config.Handler
// Capture the request Host header.
s.Server.Config.Handler = http.HandlerFunc(
func(w http.ResponseWriter, r *http.Request) {
if r.Method == "CONNECT" {
connect = true
w.WriteHeader(200)
return
}
if !connect {
t.Log("connect not recieved")
http.Error(w, "connect not recieved", 405)
return
}
origHandler.ServeHTTP(w, r)
})
ws, _, err := cstDialer.Dial(s.URL, nil)
if err != nil {
t.Fatalf("Dial: %v", err)
}
defer ws.Close()
sendRecv(t, ws)
cstDialer.Proxy = http.ProxyFromEnvironment
}
func TestProxyAuthorizationDial(t *testing.T) {
s := newServer(t)
defer s.Close()
surl, _ := url.Parse(s.URL)
surl.User = url.UserPassword("username", "password")
cstDialer.Proxy = http.ProxyURL(surl)
connect := false
origHandler := s.Server.Config.Handler
// Capture the request Host header.
s.Server.Config.Handler = http.HandlerFunc(
func(w http.ResponseWriter, r *http.Request) {
proxyAuth := r.Header.Get("Proxy-Authorization")
expectedProxyAuth := "Basic " + base64.StdEncoding.EncodeToString([]byte("username:password"))
if r.Method == "CONNECT" && proxyAuth == expectedProxyAuth {
connect = true
w.WriteHeader(200)
return
}
if !connect {
t.Log("connect with proxy authorization not recieved")
http.Error(w, "connect with proxy authorization not recieved", 405)
return
}
origHandler.ServeHTTP(w, r)
})
ws, _, err := cstDialer.Dial(s.URL, nil)
if err != nil {
t.Fatalf("Dial: %v", err)
}
defer ws.Close()
sendRecv(t, ws)
cstDialer.Proxy = http.ProxyFromEnvironment
}
func TestDial(t *testing.T) { func TestDial(t *testing.T) {
s := newServer(t) s := newServer(t)
defer s.Close() defer s.Close()
@@ -154,7 +247,7 @@ func TestDialTLS(t *testing.T) {
d := cstDialer d := cstDialer
d.NetDial = func(network, addr string) (net.Conn, error) { return net.Dial(network, u.Host) } d.NetDial = func(network, addr string) (net.Conn, error) { return net.Dial(network, u.Host) }
d.TLSClientConfig = &tls.Config{RootCAs: certs} d.TLSClientConfig = &tls.Config{RootCAs: certs}
ws, _, err := d.Dial("wss://example.com/", nil) ws, _, err := d.Dial("wss://example.com"+cstRequestURI, nil)
if err != nil { if err != nil {
t.Fatalf("Dial: %v", err) t.Fatalf("Dial: %v", err)
} }
@@ -229,6 +322,45 @@ func TestDialBadOrigin(t *testing.T) {
} }
} }
func TestDialBadHeader(t *testing.T) {
s := newServer(t)
defer s.Close()
for _, k := range []string{"Upgrade",
"Connection",
"Sec-Websocket-Key",
"Sec-Websocket-Version",
"Sec-Websocket-Protocol"} {
h := http.Header{}
h.Set(k, "bad")
ws, _, err := cstDialer.Dial(s.URL, http.Header{"Origin": {"bad"}})
if err == nil {
ws.Close()
t.Errorf("Dial with header %s returned nil", k)
}
}
}
func TestBadMethod(t *testing.T) {
s := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) {
ws, err := cstUpgrader.Upgrade(w, r, nil)
if err == nil {
t.Errorf("handshake succeeded, expect fail")
ws.Close()
}
}))
defer s.Close()
resp, err := http.PostForm(s.URL, url.Values{})
if err != nil {
t.Fatalf("PostForm returned error %v", err)
}
resp.Body.Close()
if resp.StatusCode != http.StatusMethodNotAllowed {
t.Errorf("Status = %d, want %d", resp.StatusCode, http.StatusMethodNotAllowed)
}
}
func TestHandshake(t *testing.T) { func TestHandshake(t *testing.T) {
s := newServer(t) s := newServer(t)
defer s.Close() defer s.Close()
@@ -289,8 +421,8 @@ func TestRespOnBadHandshake(t *testing.T) {
} }
} }
// If the Host header is specified in `Dial()`, the server must receive it as // TestHostHeader confirms that the host header provided in the call to Dial is
// the `Host:` header. // sent to the server.
func TestHostHeader(t *testing.T) { func TestHostHeader(t *testing.T) {
s := newServer(t) s := newServer(t)
defer s.Close() defer s.Close()
@@ -305,16 +437,12 @@ func TestHostHeader(t *testing.T) {
origHandler.ServeHTTP(w, r) origHandler.ServeHTTP(w, r)
}) })
ws, resp, err := cstDialer.Dial(s.URL, http.Header{"Host": {"testhost"}}) ws, _, err := cstDialer.Dial(s.URL, http.Header{"Host": {"testhost"}})
if err != nil { if err != nil {
t.Fatalf("Dial: %v", err) t.Fatalf("Dial: %v", err)
} }
defer ws.Close() defer ws.Close()
if resp.StatusCode != http.StatusSwitchingProtocols {
t.Fatalf("resp.StatusCode = %v, want http.StatusSwitchingProtocols", resp.StatusCode)
}
if gotHost := <-specifiedHost; gotHost != "testhost" { if gotHost := <-specifiedHost; gotHost != "testhost" {
t.Fatalf("gotHost = %q, want \"testhost\"", gotHost) t.Fatalf("gotHost = %q, want \"testhost\"", gotHost)
} }
Generated Vendored Executable → Regular
+20 -12
View File
@@ -11,16 +11,19 @@ import (
) )
var parseURLTests = []struct { var parseURLTests = []struct {
s string s string
u *url.URL u *url.URL
rui string
}{ }{
{"ws://example.com/", &url.URL{Scheme: "ws", Host: "example.com", Opaque: "/"}}, {"ws://example.com/", &url.URL{Scheme: "ws", Host: "example.com", Opaque: "/"}, "/"},
{"ws://example.com", &url.URL{Scheme: "ws", Host: "example.com", Opaque: "/"}}, {"ws://example.com", &url.URL{Scheme: "ws", Host: "example.com", Opaque: "/"}, "/"},
{"ws://example.com:7777/", &url.URL{Scheme: "ws", Host: "example.com:7777", Opaque: "/"}}, {"ws://example.com:7777/", &url.URL{Scheme: "ws", Host: "example.com:7777", Opaque: "/"}, "/"},
{"wss://example.com/", &url.URL{Scheme: "wss", Host: "example.com", Opaque: "/"}}, {"wss://example.com/", &url.URL{Scheme: "wss", Host: "example.com", Opaque: "/"}, "/"},
{"wss://example.com/a/b", &url.URL{Scheme: "wss", Host: "example.com", Opaque: "/a/b"}}, {"wss://example.com/a/b", &url.URL{Scheme: "wss", Host: "example.com", Opaque: "/a/b"}, "/a/b"},
{"ss://example.com/a/b", nil}, {"ss://example.com/a/b", nil, ""},
{"ws://[email protected]/", nil}, {"ws://[email protected]/", nil, ""},
{"wss://example.com/a/b?x=y", &url.URL{Scheme: "wss", Host: "example.com", Opaque: "/a/b", RawQuery: "x=y"}, "/a/b?x=y"},
{"wss://example.com?x=y", &url.URL{Scheme: "wss", Host: "example.com", Opaque: "/", RawQuery: "x=y"}, "/?x=y"},
} }
func TestParseURL(t *testing.T) { func TestParseURL(t *testing.T) {
@@ -30,14 +33,19 @@ func TestParseURL(t *testing.T) {
t.Errorf("parseURL(%q) returned error %v", tt.s, err) t.Errorf("parseURL(%q) returned error %v", tt.s, err)
continue continue
} }
if tt.u == nil && err == nil { if tt.u == nil {
t.Errorf("parseURL(%q) did not return error", tt.s) if err == nil {
t.Errorf("parseURL(%q) did not return error", tt.s)
}
continue continue
} }
if !reflect.DeepEqual(u, tt.u) { if !reflect.DeepEqual(u, tt.u) {
t.Errorf("parseURL(%q) returned %v, want %v", tt.s, u, tt.u) t.Errorf("parseURL(%q) = %v, want %v", tt.s, u, tt.u)
continue continue
} }
if u.RequestURI() != tt.rui {
t.Errorf("parseURL(%q).RequestURI() = %v, want %v", tt.s, u.RequestURI(), tt.rui)
}
} }
} }
Generated Vendored Executable → Regular
+118 -28
View File
@@ -88,19 +88,82 @@ func (e *netError) Error() string { return e.msg }
func (e *netError) Temporary() bool { return e.temporary } func (e *netError) Temporary() bool { return e.temporary }
func (e *netError) Timeout() bool { return e.timeout } func (e *netError) Timeout() bool { return e.timeout }
// closeError represents close frame. // CloseError represents close frame.
type closeError struct { type CloseError struct {
code int
text string // Code is defined in RFC 6455, section 11.7.
Code int
// Text is the optional text payload.
Text string
} }
func (e *closeError) Error() string { func (e *CloseError) Error() string {
return "websocket: close " + strconv.Itoa(e.code) + " " + e.text s := []byte("websocket: close ")
s = strconv.AppendInt(s, int64(e.Code), 10)
switch e.Code {
case CloseNormalClosure:
s = append(s, " (normal)"...)
case CloseGoingAway:
s = append(s, " (going away)"...)
case CloseProtocolError:
s = append(s, " (protocol error)"...)
case CloseUnsupportedData:
s = append(s, " (unsupported data)"...)
case CloseNoStatusReceived:
s = append(s, " (no status)"...)
case CloseAbnormalClosure:
s = append(s, " (abnormal closure)"...)
case CloseInvalidFramePayloadData:
s = append(s, " (invalid payload data)"...)
case ClosePolicyViolation:
s = append(s, " (policy violation)"...)
case CloseMessageTooBig:
s = append(s, " (message too big)"...)
case CloseMandatoryExtension:
s = append(s, " (mandatory extension missing)"...)
case CloseInternalServerErr:
s = append(s, " (internal server error)"...)
case CloseTLSHandshake:
s = append(s, " (TLS handshake error)"...)
}
if e.Text != "" {
s = append(s, ": "...)
s = append(s, e.Text...)
}
return string(s)
}
// IsCloseError returns boolean indicating whether the error is a *CloseError
// with one of the specified codes.
func IsCloseError(err error, codes ...int) bool {
if e, ok := err.(*CloseError); ok {
for _, code := range codes {
if e.Code == code {
return true
}
}
}
return false
}
// IsUnexpectedCloseError returns boolean indicating whether the error is a
// *CloseError with a code not in the list of expected codes.
func IsUnexpectedCloseError(err error, expectedCodes ...int) bool {
if e, ok := err.(*CloseError); ok {
for _, code := range expectedCodes {
if e.Code == code {
return false
}
}
return true
}
return false
} }
var ( var (
errWriteTimeout = &netError{msg: "websocket: write timeout", timeout: true} errWriteTimeout = &netError{msg: "websocket: write timeout", timeout: true, temporary: true}
errUnexpectedEOF = &closeError{code: CloseAbnormalClosure, text: io.ErrUnexpectedEOF.Error()} errUnexpectedEOF = &CloseError{Code: CloseAbnormalClosure, Text: io.ErrUnexpectedEOF.Error()}
errBadWriteOpCode = errors.New("websocket: bad write message type") errBadWriteOpCode = errors.New("websocket: bad write message type")
errWriteClosed = errors.New("websocket: write closed") errWriteClosed = errors.New("websocket: write closed")
errInvalidControlFrame = errors.New("websocket: invalid control frame") errInvalidControlFrame = errors.New("websocket: invalid control frame")
@@ -151,6 +214,7 @@ type Conn struct {
writeFrameType int // type of the current frame. writeFrameType int // type of the current frame.
writeSeq int // incremented to invalidate message writers. writeSeq int // incremented to invalidate message writers.
writeDeadline time.Time writeDeadline time.Time
isWriting bool // for best-effort concurrent write detection
// Read fields // Read fields
readErr error readErr error
@@ -164,6 +228,7 @@ type Conn struct {
readMaskKey [4]byte readMaskKey [4]byte
handlePong func(string) error handlePong func(string) error
handlePing func(string) error handlePing func(string) error
readErrCount int
} }
func newConn(conn net.Conn, isServer bool, readBufferSize, writeBufferSize int) *Conn { func newConn(conn net.Conn, isServer bool, readBufferSize, writeBufferSize int) *Conn {
@@ -296,7 +361,7 @@ func (c *Conn) WriteControl(messageType int, data []byte, deadline time.Time) er
if n != 0 && n != len(buf) { if n != 0 && n != len(buf) {
c.conn.Close() c.conn.Close()
} }
return err return hideTempErr(err)
} }
// NextWriter returns a writer for the next message to send. The writer's // NextWriter returns a writer for the next message to send. The writer's
@@ -304,9 +369,6 @@ func (c *Conn) WriteControl(messageType int, data []byte, deadline time.Time) er
// //
// There can be at most one open writer on a connection. NextWriter closes the // There can be at most one open writer on a connection. NextWriter closes the
// previous writer if the application has not already done so. // previous writer if the application has not already done so.
//
// The NextWriter method and the writers returned from the method cannot be
// accessed by more than one goroutine at a time.
func (c *Conn) NextWriter(messageType int) (io.WriteCloser, error) { func (c *Conn) NextWriter(messageType int) (io.WriteCloser, error) {
if c.writeErr != nil { if c.writeErr != nil {
return nil, c.writeErr return nil, c.writeErr
@@ -380,9 +442,22 @@ func (c *Conn) flushFrame(final bool, extra []byte) error {
} }
} }
// Write the buffers to the connection. // Write the buffers to the connection with best-effort detection of
// concurrent writes. See the concurrency section in the package
// documentation for more info.
if c.isWriting {
panic("concurrent write to websocket connection")
}
c.isWriting = true
c.writeErr = c.write(c.writeFrameType, c.writeDeadline, c.writeBuf[framePos:c.writePos], extra) c.writeErr = c.write(c.writeFrameType, c.writeDeadline, c.writeBuf[framePos:c.writePos], extra)
if !c.isWriting {
panic("concurrent write to websocket connection")
}
c.isWriting = false
// Setup for next frame. // Setup for next frame.
c.writePos = maxFrameHeaderSize c.writePos = maxFrameHeaderSize
c.writeFrameType = continuationFrame c.writeFrameType = continuationFrame
@@ -666,19 +741,16 @@ func (c *Conn) advanceFrame() (int, error) {
return noFrame, err return noFrame, err
} }
case CloseMessage: case CloseMessage:
c.WriteControl(CloseMessage, []byte{}, time.Now().Add(writeWait)) echoMessage := []byte{}
closeCode := CloseNoStatusReceived closeCode := CloseNoStatusReceived
closeText := "" closeText := ""
if len(payload) >= 2 { if len(payload) >= 2 {
echoMessage = payload[:2]
closeCode = int(binary.BigEndian.Uint16(payload)) closeCode = int(binary.BigEndian.Uint16(payload))
closeText = string(payload[2:]) closeText = string(payload[2:])
} }
switch closeCode { c.WriteControl(CloseMessage, echoMessage, time.Now().Add(writeWait))
case CloseNormalClosure, CloseGoingAway: return noFrame, &CloseError{Code: closeCode, Text: closeText}
return noFrame, io.EOF
default:
return noFrame, &closeError{code: closeCode, text: closeText}
}
} }
return frameType, nil return frameType, nil
@@ -695,8 +767,10 @@ func (c *Conn) handleProtocolError(message string) error {
// There can be at most one open reader on a connection. NextReader discards // There can be at most one open reader on a connection. NextReader discards
// the previous message if the application has not already consumed it. // the previous message if the application has not already consumed it.
// //
// The NextReader method and the readers returned from the method cannot be // Applications must break out of the application's read loop when this method
// accessed by more than one goroutine at a time. // returns a non-nil error value. Errors returned from this method are
// permanent. Once this method returns a non-nil error, all subsequent calls to
// this method return the same error.
func (c *Conn) NextReader() (messageType int, r io.Reader, err error) { func (c *Conn) NextReader() (messageType int, r io.Reader, err error) {
c.readSeq++ c.readSeq++
@@ -712,6 +786,15 @@ func (c *Conn) NextReader() (messageType int, r io.Reader, err error) {
return frameType, messageReader{c, c.readSeq}, nil return frameType, messageReader{c, c.readSeq}, nil
} }
} }
// Applications that do handle the error returned from this method spin in
// tight loop on connection failure. To help application developers detect
// this error, panic on repeated reads to the failed connection.
c.readErrCount++
if c.readErrCount >= 1000 {
panic("repeated read on failed websocket connection")
}
return noFrame, nil, c.readErr return noFrame, nil, c.readErr
} }
@@ -790,20 +873,27 @@ func (c *Conn) SetReadLimit(limit int64) {
} }
// SetPingHandler sets the handler for ping messages received from the peer. // SetPingHandler sets the handler for ping messages received from the peer.
// The default ping handler sends a pong to the peer. // The appData argument to h is the PING frame application data. The default
func (c *Conn) SetPingHandler(h func(string) error) { // ping handler sends a pong to the peer.
func (c *Conn) SetPingHandler(h func(appData string) error) {
if h == nil { if h == nil {
h = func(message string) error { h = func(message string) error {
c.WriteControl(PongMessage, []byte(message), time.Now().Add(writeWait)) err := c.WriteControl(PongMessage, []byte(message), time.Now().Add(writeWait))
return nil if err == ErrCloseSent {
return nil
} else if e, ok := err.(net.Error); ok && e.Temporary() {
return nil
}
return err
} }
} }
c.handlePing = h c.handlePing = h
} }
// SetPongHandler sets the handler for pong messages received from the peer. // SetPongHandler sets the handler for pong messages received from the peer.
// The default pong handler does nothing. // The appData argument to h is the PONG frame application data. The default
func (c *Conn) SetPongHandler(h func(string) error) { // pong handler does nothing.
func (c *Conn) SetPongHandler(h func(appData string) error) {
if h == nil { if h == nil {
h = func(string) error { return nil } h = func(string) error { return nil }
} }
Generated Vendored Executable → Regular
+134 -5
View File
@@ -5,11 +5,14 @@
package websocket package websocket
import ( import (
"bufio"
"bytes" "bytes"
"errors"
"fmt" "fmt"
"io" "io"
"io/ioutil" "io/ioutil"
"net" "net"
"reflect"
"testing" "testing"
"testing/iotest" "testing/iotest"
"time" "time"
@@ -146,13 +149,15 @@ func TestControl(t *testing.T) {
func TestCloseBeforeFinalFrame(t *testing.T) { func TestCloseBeforeFinalFrame(t *testing.T) {
const bufSize = 512 const bufSize = 512
expectedErr := &CloseError{Code: CloseNormalClosure, Text: "hello"}
var b1, b2 bytes.Buffer var b1, b2 bytes.Buffer
wc := newConn(fakeNetConn{Reader: nil, Writer: &b1}, false, 1024, bufSize) wc := newConn(fakeNetConn{Reader: nil, Writer: &b1}, false, 1024, bufSize)
rc := newConn(fakeNetConn{Reader: &b1, Writer: &b2}, true, 1024, 1024) rc := newConn(fakeNetConn{Reader: &b1, Writer: &b2}, true, 1024, 1024)
w, _ := wc.NextWriter(BinaryMessage) w, _ := wc.NextWriter(BinaryMessage)
w.Write(make([]byte, bufSize+bufSize/2)) w.Write(make([]byte, bufSize+bufSize/2))
wc.WriteControl(CloseMessage, FormatCloseMessage(CloseNormalClosure, ""), time.Now().Add(10*time.Second)) wc.WriteControl(CloseMessage, FormatCloseMessage(expectedErr.Code, expectedErr.Text), time.Now().Add(10*time.Second))
w.Close() w.Close()
op, r, err := rc.NextReader() op, r, err := rc.NextReader()
@@ -160,12 +165,12 @@ func TestCloseBeforeFinalFrame(t *testing.T) {
t.Fatalf("NextReader() returned %d, %v", op, err) t.Fatalf("NextReader() returned %d, %v", op, err)
} }
_, err = io.Copy(ioutil.Discard, r) _, err = io.Copy(ioutil.Discard, r)
if err != errUnexpectedEOF { if !reflect.DeepEqual(err, expectedErr) {
t.Fatalf("io.Copy() returned %v, want %v", err, errUnexpectedEOF) t.Fatalf("io.Copy() returned %v, want %v", err, expectedErr)
} }
_, _, err = rc.NextReader() _, _, err = rc.NextReader()
if err != io.EOF { if !reflect.DeepEqual(err, expectedErr) {
t.Fatalf("NextReader() returned %v, want %v", err, io.EOF) t.Fatalf("NextReader() returned %v, want %v", err, expectedErr)
} }
} }
@@ -236,3 +241,127 @@ func TestUnderlyingConn(t *testing.T) {
t.Fatalf("Underlying conn is not what it should be.") t.Fatalf("Underlying conn is not what it should be.")
} }
} }
func TestBufioReadBytes(t *testing.T) {
// Test calling bufio.ReadBytes for value longer than read buffer size.
m := make([]byte, 512)
m[len(m)-1] = '\n'
var b1, b2 bytes.Buffer
wc := newConn(fakeNetConn{Reader: nil, Writer: &b1}, false, len(m)+64, len(m)+64)
rc := newConn(fakeNetConn{Reader: &b1, Writer: &b2}, true, len(m)-64, len(m)-64)
w, _ := wc.NextWriter(BinaryMessage)
w.Write(m)
w.Close()
op, r, err := rc.NextReader()
if op != BinaryMessage || err != nil {
t.Fatalf("NextReader() returned %d, %v", op, err)
}
br := bufio.NewReader(r)
p, err := br.ReadBytes('\n')
if err != nil {
t.Fatalf("ReadBytes() returned %v", err)
}
if len(p) != len(m) {
t.Fatalf("read returnd %d bytes, want %d bytes", len(p), len(m))
}
}
var closeErrorTests = []struct {
err error
codes []int
ok bool
}{
{&CloseError{Code: CloseNormalClosure}, []int{CloseNormalClosure}, true},
{&CloseError{Code: CloseNormalClosure}, []int{CloseNoStatusReceived}, false},
{&CloseError{Code: CloseNormalClosure}, []int{CloseNoStatusReceived, CloseNormalClosure}, true},
{errors.New("hello"), []int{CloseNormalClosure}, false},
}
func TestCloseError(t *testing.T) {
for _, tt := range closeErrorTests {
ok := IsCloseError(tt.err, tt.codes...)
if ok != tt.ok {
t.Errorf("IsCloseError(%#v, %#v) returned %v, want %v", tt.err, tt.codes, ok, tt.ok)
}
}
}
var unexpectedCloseErrorTests = []struct {
err error
codes []int
ok bool
}{
{&CloseError{Code: CloseNormalClosure}, []int{CloseNormalClosure}, false},
{&CloseError{Code: CloseNormalClosure}, []int{CloseNoStatusReceived}, true},
{&CloseError{Code: CloseNormalClosure}, []int{CloseNoStatusReceived, CloseNormalClosure}, false},
{errors.New("hello"), []int{CloseNormalClosure}, false},
}
func TestUnexpectedCloseErrors(t *testing.T) {
for _, tt := range unexpectedCloseErrorTests {
ok := IsUnexpectedCloseError(tt.err, tt.codes...)
if ok != tt.ok {
t.Errorf("IsUnexpectedCloseError(%#v, %#v) returned %v, want %v", tt.err, tt.codes, ok, tt.ok)
}
}
}
type blockingWriter struct {
c1, c2 chan struct{}
}
func (w blockingWriter) Write(p []byte) (int, error) {
// Allow main to continue
close(w.c1)
// Wait for panic in main
<-w.c2
return len(p), nil
}
func TestConcurrentWritePanic(t *testing.T) {
w := blockingWriter{make(chan struct{}), make(chan struct{})}
c := newConn(fakeNetConn{Reader: nil, Writer: w}, false, 1024, 1024)
go func() {
c.WriteMessage(TextMessage, []byte{})
}()
// wait for goroutine to block in write.
<-w.c1
defer func() {
close(w.c2)
if v := recover(); v != nil {
return
}
}()
c.WriteMessage(TextMessage, []byte{})
t.Fatal("should not get here")
}
type failingReader struct{}
func (r failingReader) Read(p []byte) (int, error) {
return 0, io.EOF
}
func TestFailedConnectionReadPanic(t *testing.T) {
c := newConn(fakeNetConn{Reader: failingReader{}, Writer: nil}, false, 1024, 1024)
defer func() {
if v := recover(); v != nil {
return
}
}()
for i := 0; i < 20000; i++ {
c.ReadMessage()
}
t.Fatal("should not get here")
}
Generated Vendored Executable → Regular
+25 -25
View File
@@ -46,8 +46,7 @@
// method to get an io.WriteCloser, write the message to the writer and close // method to get an io.WriteCloser, write the message to the writer and close
// the writer when done. To receive a message, call the connection NextReader // the writer when done. To receive a message, call the connection NextReader
// method to get an io.Reader and read until io.EOF is returned. This snippet // method to get an io.Reader and read until io.EOF is returned. This snippet
// snippet shows how to echo messages using the NextWriter and NextReader // shows how to echo messages using the NextWriter and NextReader methods:
// methods:
// //
// for { // for {
// messageType, r, err := conn.NextReader() // messageType, r, err := conn.NextReader()
@@ -86,31 +85,19 @@
// and pong. Call the connection WriteControl, WriteMessage or NextWriter // and pong. Call the connection WriteControl, WriteMessage or NextWriter
// methods to send a control message to the peer. // methods to send a control message to the peer.
// //
// Connections handle received ping and pong messages by invoking a callback // Connections handle received ping and pong messages by invoking callback
// function set with SetPingHandler and SetPongHandler methods. These callback // functions set with SetPingHandler and SetPongHandler methods. The default
// functions can be invoked from the ReadMessage method, the NextReader method // ping handler sends a pong to the client. The callback functions can be
// or from a call to the data message reader returned from NextReader. // invoked from the NextReader, ReadMessage or the message Read method.
// //
// Connections handle received close messages by returning an error from the // Connections handle received close messages by sending a close message to the
// ReadMessage method, the NextReader method or from a call to the data message // peer and returning a *CloseError from the the NextReader, ReadMessage or the
// reader returned from NextReader. // message Read method.
//
// Concurrency
//
// Connections do not support concurrent calls to the write methods
// (NextWriter, SetWriteDeadline, WriteMessage) or concurrent calls to the read
// methods methods (NextReader, SetReadDeadline, ReadMessage). Connections do
// support a concurrent reader and writer.
//
// The Close and WriteControl methods can be called concurrently with all other
// methods.
//
// Read is Required
// //
// The application must read the connection to process ping and close messages // The application must read the connection to process ping and close messages
// sent from the peer. If the application is not otherwise interested in // sent from the peer. If the application is not otherwise interested in
// messages from the peer, then the application should start a goroutine to read // messages from the peer, then the application should start a goroutine to
// and discard messages from the peer. A simple example is: // read and discard messages from the peer. A simple example is:
// //
// func readLoop(c *websocket.Conn) { // func readLoop(c *websocket.Conn) {
// for { // for {
@@ -121,6 +108,19 @@
// } // }
// } // }
// //
// Concurrency
//
// Connections support one concurrent reader and one concurrent writer.
//
// Applications are responsible for ensuring that no more than one goroutine
// calls the write methods (NextWriter, SetWriteDeadline, WriteMessage,
// WriteJSON) concurrently and that no more than one goroutine calls the read
// methods (NextReader, SetReadDeadline, ReadMessage, ReadJSON, SetPongHandler,
// SetPingHandler) concurrently.
//
// The Close and WriteControl methods can be called concurrently with all other
// methods.
//
// Origin Considerations // Origin Considerations
// //
// Web browsers allow Javascript applications to open a WebSocket connection to // Web browsers allow Javascript applications to open a WebSocket connection to
@@ -138,9 +138,9 @@
// An application can allow connections from any origin by specifying a // An application can allow connections from any origin by specifying a
// function that always returns true: // function that always returns true:
// //
// var upgrader = websocket.Upgrader{ // var upgrader = websocket.Upgrader{
// CheckOrigin: func(r *http.Request) bool { return true }, // CheckOrigin: func(r *http.Request) bool { return true },
// } // }
// //
// The deprecated Upgrade function does not enforce an origin policy. It's the // The deprecated Upgrade function does not enforce an origin policy. It's the
// application's responsibility to check the Origin header before calling // application's responsibility to check the Origin header before calling
+40
View File
@@ -0,0 +1,40 @@
// Copyright 2015 The Gorilla WebSocket Authors. All rights reserved.
// Use of this source code is governed by a BSD-style
// license that can be found in the LICENSE file.
package websocket_test
import (
"log"
"net/http"
"testing"
"github.com/gorilla/websocket"
)
// The websocket.IsUnexpectedCloseError function is useful for identifying
// application and protocol errors.
//
// This server application works with a client application running in the
// browser. The client application does not explicitly close the websocket. The
// only expected close message from the client has the code
// websocket.CloseGoingAway. All other other close messages are likely the
// result of an application or protocol error and are logged to aid debugging.
func ExampleIsUnexpectedCloseError(err error, c *websocket.Conn, req *http.Request) {
for {
messageType, p, err := c.ReadMessage()
if err != nil {
if websocket.IsUnexpectedCloseError(err, websocket.CloseGoingAway) {
log.Printf("error: %v, user-agent: %v", err, req.Header.Get("User-Agent"))
}
return
}
processMesage(messageType, p)
}
}
func processMesage(mt int, p []byte) {}
// TestX prevents godoc from showing this entire file in the example. Remove
// this function when a second example is added.
func TestX(t *testing.T) {}
View File
View File
View File
Generated Vendored Executable → Regular
+1
View File
@@ -17,3 +17,4 @@ using the following commands.
$ cd `go list -f '{{.Dir}}' github.com/gorilla/websocket/examples/chat` $ cd `go list -f '{{.Dir}}' github.com/gorilla/websocket/examples/chat`
$ go run *.go $ go run *.go
To use the chat example, open http://localhost:8080/ in your browser.
Generated Vendored Executable → Regular
+4 -5
View File
@@ -51,6 +51,9 @@ func (c *connection) readPump() {
for { for {
_, message, err := c.ws.ReadMessage() _, message, err := c.ws.ReadMessage()
if err != nil { if err != nil {
if websocket.IsUnexpectedCloseError(err, websocket.CloseGoingAway) {
log.Printf("error: %v", err)
}
break break
} }
h.broadcast <- message h.broadcast <- message
@@ -88,12 +91,8 @@ func (c *connection) writePump() {
} }
} }
// serverWs handles websocket requests from the peer. // serveWs handles websocket requests from the peer.
func serveWs(w http.ResponseWriter, r *http.Request) { func serveWs(w http.ResponseWriter, r *http.Request) {
if r.Method != "GET" {
http.Error(w, "Method not allowed", 405)
return
}
ws, err := upgrader.Upgrade(w, r, nil) ws, err := upgrader.Upgrade(w, r, nil)
if err != nil { if err != nil {
log.Println(err) log.Println(err)
Generated Vendored Executable → Regular
View File
Generated Vendored Executable → Regular
View File
Generated Vendored Executable → Regular
View File
+19
View File
@@ -0,0 +1,19 @@
# Command example
This example connects a websocket connection to stdin and stdout of a command.
Received messages are written to stdin followed by a `\n`. Each line read from
from standard out is sent as a message to the client.
$ go get github.com/gorilla/websocket
$ cd `go list -f '{{.Dir}}' github.com/gorilla/websocket/examples/command`
$ go run main.go <command and arguments to run>
# Open http://localhost:8080/ .
Try the following commands.
# Echo sent messages to the output area.
$ go run main.go cat
# Run a shell.Try sending "ls" and "cat main.go".
$ go run main.go sh
+96
View File
@@ -0,0 +1,96 @@
<!DOCTYPE html>
<html lang="en">
<head>
<title>Command Example</title>
<script src="//ajax.googleapis.com/ajax/libs/jquery/2.0.3/jquery.min.js"></script>
<script type="text/javascript">
$(function() {
var conn;
var msg = $("#msg");
var log = $("#log");
function appendLog(msg) {
var d = log[0]
var doScroll = d.scrollTop == d.scrollHeight - d.clientHeight;
msg.appendTo(log)
if (doScroll) {
d.scrollTop = d.scrollHeight - d.clientHeight;
}
}
$("#form").submit(function() {
if (!conn) {
return false;
}
if (!msg.val()) {
return false;
}
conn.send(msg.val());
msg.val("");
return false
});
if (window["WebSocket"]) {
conn = new WebSocket("ws://{{$}}/ws");
conn.onclose = function(evt) {
appendLog($("<div><b>Connection closed.</b></div>"))
}
conn.onmessage = function(evt) {
appendLog($("<pre/>").text(evt.data))
}
} else {
appendLog($("<div><b>Your browser does not support WebSockets.</b></div>"))
}
});
</script>
<style type="text/css">
html {
overflow: hidden;
}
body {
overflow: hidden;
padding: 0;
margin: 0;
width: 100%;
height: 100%;
background: gray;
}
#log {
background: white;
margin: 0;
padding: 0.5em 0.5em 0.5em 0.5em;
position: absolute;
top: 0.5em;
left: 0.5em;
right: 0.5em;
bottom: 3em;
overflow: auto;
}
#log pre {
margin: 0;
}
#form {
padding: 0 0.5em 0 0.5em;
margin: 0;
position: absolute;
bottom: 1em;
left: 0px;
width: 100%;
overflow: hidden;
}
</style>
</head>
<body>
<div id="log"></div>
<form id="form">
<input type="submit" value="Send" />
<input type="text" id="msg" size="64"/>
</form>
</body>
</html>
+188
View File
@@ -0,0 +1,188 @@
// Copyright 2015 The Gorilla WebSocket Authors. All rights reserved.
// Use of this source code is governed by a BSD-style
// license that can be found in the LICENSE file.
package main
import (
"bufio"
"flag"
"io"
"log"
"net/http"
"os"
"os/exec"
"text/template"
"time"
"github.com/gorilla/websocket"
)
var (
addr = flag.String("addr", "127.0.0.1:8080", "http service address")
cmdPath string
homeTempl = template.Must(template.ParseFiles("home.html"))
)
const (
// Time allowed to write a message to the peer.
writeWait = 10 * time.Second
// Maximum message size allowed from peer.
maxMessageSize = 8192
// Time allowed to read the next pong message from the peer.
pongWait = 60 * time.Second
// Send pings to peer with this period. Must be less than pongWait.
pingPeriod = (pongWait * 9) / 10
)
func pumpStdin(ws *websocket.Conn, w io.Writer) {
defer ws.Close()
ws.SetReadLimit(maxMessageSize)
ws.SetReadDeadline(time.Now().Add(pongWait))
ws.SetPongHandler(func(string) error { ws.SetReadDeadline(time.Now().Add(pongWait)); return nil })
for {
_, message, err := ws.ReadMessage()
if err != nil {
break
}
message = append(message, '\n')
if _, err := w.Write(message); err != nil {
break
}
}
}
func pumpStdout(ws *websocket.Conn, r io.Reader, done chan struct{}) {
defer func() {
ws.Close()
close(done)
}()
s := bufio.NewScanner(r)
for s.Scan() {
ws.SetWriteDeadline(time.Now().Add(writeWait))
if err := ws.WriteMessage(websocket.TextMessage, s.Bytes()); err != nil {
break
}
}
if s.Err() != nil {
log.Println("scan:", s.Err())
}
}
func ping(ws *websocket.Conn, done chan struct{}) {
ticker := time.NewTicker(pingPeriod)
defer ticker.Stop()
for {
select {
case <-ticker.C:
if err := ws.WriteControl(websocket.PingMessage, []byte{}, time.Now().Add(writeWait)); err != nil {
log.Println("ping:", err)
}
case <-done:
return
}
}
}
func internalError(ws *websocket.Conn, msg string, err error) {
log.Println(msg, err)
ws.WriteMessage(websocket.TextMessage, []byte("Internal server error."))
}
var upgrader = websocket.Upgrader{}
func serveWs(w http.ResponseWriter, r *http.Request) {
ws, err := upgrader.Upgrade(w, r, nil)
if err != nil {
log.Println("upgrade:", err)
return
}
defer ws.Close()
outr, outw, err := os.Pipe()
if err != nil {
internalError(ws, "stdout:", err)
return
}
defer outr.Close()
defer outw.Close()
inr, inw, err := os.Pipe()
if err != nil {
internalError(ws, "stdin:", err)
return
}
defer inr.Close()
defer inw.Close()
proc, err := os.StartProcess(cmdPath, flag.Args(), &os.ProcAttr{
Files: []*os.File{inr, outw, outw},
})
if err != nil {
internalError(ws, "start:", err)
return
}
inr.Close()
outw.Close()
stdoutDone := make(chan struct{})
go pumpStdout(ws, outr, stdoutDone)
go ping(ws, stdoutDone)
pumpStdin(ws, inw)
// Some commands will exit when stdin is closed.
inw.Close()
// Other commands need a bonk on the head.
if err := proc.Signal(os.Interrupt); err != nil {
log.Println("inter:", err)
}
select {
case <-stdoutDone:
case <-time.After(time.Second):
// A bigger bonk on the head.
if err := proc.Signal(os.Kill); err != nil {
log.Println("term:", err)
}
<-stdoutDone
}
if _, err := proc.Wait(); err != nil {
log.Println("wait:", err)
}
}
func serveHome(w http.ResponseWriter, r *http.Request) {
if r.URL.Path != "/" {
http.Error(w, "Not found", 404)
return
}
if r.Method != "GET" {
http.Error(w, "Method not allowed", 405)
return
}
w.Header().Set("Content-Type", "text/html; charset=utf-8")
homeTempl.Execute(w, r.Host)
}
func main() {
flag.Parse()
if len(flag.Args()) < 1 {
log.Fatal("must specify at least one argument")
}
var err error
cmdPath, err = exec.LookPath(flag.Args()[0])
if err != nil {
log.Fatal(err)
}
http.HandleFunc("/", serveHome)
http.HandleFunc("/ws", serveWs)
log.Fatal(http.ListenAndServe(*addr, nil))
}
+17
View File
@@ -0,0 +1,17 @@
# Client and server example
This example shows a simple client and server.
The server echoes messages sent to it. The client sends a message every second
and prints all messages received.
To run the example, start the server:
$ go run server.go
Next, start the client:
$ go run client.go
The server includes a simple web client. To use the client, open
http://127.0.0.1:8080 in the browser and follow the instructions on the page.
+81
View File
@@ -0,0 +1,81 @@
// Copyright 2015 The Gorilla WebSocket Authors. All rights reserved.
// Use of this source code is governed by a BSD-style
// license that can be found in the LICENSE file.
// +build ignore
package main
import (
"flag"
"log"
"net/url"
"os"
"os/signal"
"time"
"github.com/gorilla/websocket"
)
var addr = flag.String("addr", "localhost:8080", "http service address")
func main() {
flag.Parse()
log.SetFlags(0)
interrupt := make(chan os.Signal, 1)
signal.Notify(interrupt, os.Interrupt)
u := url.URL{Scheme: "ws", Host: *addr, Path: "/echo"}
log.Printf("connecting to %s", u.String())
c, _, err := websocket.DefaultDialer.Dial(u.String(), nil)
if err != nil {
log.Fatal("dial:", err)
}
defer c.Close()
done := make(chan struct{})
go func() {
defer c.Close()
defer close(done)
for {
_, message, err := c.ReadMessage()
if err != nil {
log.Println("read:", err)
return
}
log.Printf("recv: %s", message)
}
}()
ticker := time.NewTicker(time.Second)
defer ticker.Stop()
for {
select {
case t := <-ticker.C:
err := c.WriteMessage(websocket.TextMessage, []byte(t.String()))
if err != nil {
log.Println("write:", err)
return
}
case <-interrupt:
log.Println("interrupt")
// To cleanly close a connection, a client should send a close
// frame and wait for the server to close the connection.
err := c.WriteMessage(websocket.CloseMessage, websocket.FormatCloseMessage(websocket.CloseNormalClosure, ""))
if err != nil {
log.Println("write close:", err)
return
}
select {
case <-done:
case <-time.After(time.Second):
}
c.Close()
return
}
}
}
+132
View File
@@ -0,0 +1,132 @@
// Copyright 2015 The Gorilla WebSocket Authors. All rights reserved.
// Use of this source code is governed by a BSD-style
// license that can be found in the LICENSE file.
// +build ignore
package main
import (
"flag"
"html/template"
"log"
"net/http"
"github.com/gorilla/websocket"
)
var addr = flag.String("addr", "localhost:8080", "http service address")
var upgrader = websocket.Upgrader{} // use default options
func echo(w http.ResponseWriter, r *http.Request) {
c, err := upgrader.Upgrade(w, r, nil)
if err != nil {
log.Print("upgrade:", err)
return
}
defer c.Close()
for {
mt, message, err := c.ReadMessage()
if err != nil {
log.Println("read:", err)
break
}
log.Printf("recv: %s", message)
err = c.WriteMessage(mt, message)
if err != nil {
log.Println("write:", err)
break
}
}
}
func home(w http.ResponseWriter, r *http.Request) {
homeTemplate.Execute(w, "ws://"+r.Host+"/echo")
}
func main() {
flag.Parse()
log.SetFlags(0)
http.HandleFunc("/echo", echo)
http.HandleFunc("/", home)
log.Fatal(http.ListenAndServe(*addr, nil))
}
var homeTemplate = template.Must(template.New("").Parse(`
<!DOCTYPE html>
<head>
<meta charset="utf-8">
<script>
window.addEventListener("load", function(evt) {
var output = document.getElementById("output");
var input = document.getElementById("input");
var ws;
var print = function(message) {
var d = document.createElement("div");
d.innerHTML = message;
output.appendChild(d);
};
document.getElementById("open").onclick = function(evt) {
if (ws) {
return false;
}
ws = new WebSocket("{{.}}");
ws.onopen = function(evt) {
print("OPEN");
}
ws.onclose = function(evt) {
print("CLOSE");
ws = null;
}
ws.onmessage = function(evt) {
print("RESPONSE: " + evt.data);
}
ws.onerror = function(evt) {
print("ERROR: " + evt.data);
}
return false;
};
document.getElementById("send").onclick = function(evt) {
if (!ws) {
return false;
}
print("SEND: " + input.value);
ws.send(input.value);
return false;
};
document.getElementById("close").onclick = function(evt) {
if (!ws) {
return false;
}
ws.close();
return false;
};
});
</script>
</head>
<body>
<table>
<tr><td valign="top" width="50%">
<p>Click "Open" to create a connection to the server,
"Send" to send a message to the server and "Close" to close the connection.
You can change the message and send multiple times.
<p>
<form>
<button id="open">Open</button>
<button id="close">Close</button>
<p><input id="input" type="text" value="Hello world!">
<button id="send">Send</button>
</form>
</td><td valign="top" width="50%">
<div id="output"></div>
</td></tr></table>
</body>
</html>
`))
View File
View File
Generated Vendored Executable → Regular
+1 -3
View File
@@ -48,9 +48,7 @@ func (c *Conn) ReadJSON(v interface{}) error {
} }
err = json.NewDecoder(r).Decode(v) err = json.NewDecoder(r).Decode(v)
if err == io.EOF { if err == io.EOF {
// Decode returns io.EOF when the message is empty or all whitespace. // One value is expected in the message.
// Convert to io.ErrUnexpectedEOF so that application can distinguish
// between an error reading the JSON value and the connection closing.
err = io.ErrUnexpectedEOF err = io.ErrUnexpectedEOF
} }
return err return err
Generated Vendored Executable → Regular
+2 -2
View File
@@ -38,7 +38,7 @@ func TestJSON(t *testing.T) {
} }
} }
func TestPartialJsonRead(t *testing.T) { func TestPartialJSONRead(t *testing.T) {
var buf bytes.Buffer var buf bytes.Buffer
c := fakeNetConn{&buf, &buf} c := fakeNetConn{&buf, &buf}
wc := newConn(c, true, 1024, 1024) wc := newConn(c, true, 1024, 1024)
@@ -87,7 +87,7 @@ func TestPartialJsonRead(t *testing.T) {
} }
err = rc.ReadJSON(&v) err = rc.ReadJSON(&v)
if err != io.EOF { if _, ok := err.(*CloseError); !ok {
t.Error("final", err) t.Error("final", err)
} }
} }
Generated Vendored Executable → Regular
+6
View File
@@ -92,7 +92,13 @@ func (u *Upgrader) selectSubprotocol(r *http.Request, responseHeader http.Header
// The responseHeader is included in the response to the client's upgrade // The responseHeader is included in the response to the client's upgrade
// request. Use the responseHeader to specify cookies (Set-Cookie) and the // request. Use the responseHeader to specify cookies (Set-Cookie) and the
// application negotiated subprotocol (Sec-Websocket-Protocol). // application negotiated subprotocol (Sec-Websocket-Protocol).
//
// If the upgrade fails, then Upgrade replies to the client with an HTTP error
// response.
func (u *Upgrader) Upgrade(w http.ResponseWriter, r *http.Request, responseHeader http.Header) (*Conn, error) { func (u *Upgrader) Upgrade(w http.ResponseWriter, r *http.Request, responseHeader http.Header) (*Conn, error) {
if r.Method != "GET" {
return u.returnError(w, r, http.StatusMethodNotAllowed, "websocket: method not GET")
}
if values := r.Header["Sec-Websocket-Version"]; len(values) == 0 || values[0] != "13" { if values := r.Header["Sec-Websocket-Version"]; len(values) == 0 || values[0] != "13" {
return u.returnError(w, r, http.StatusBadRequest, "websocket: version != 13") return u.returnError(w, r, http.StatusBadRequest, "websocket: version != 13")
} }
Generated Vendored Executable → Regular
View File
Generated Vendored Executable → Regular
View File
Generated Vendored Executable → Regular
View File
Generated Vendored Submodule
+1
Submodule vendor/github.com/pusher/pusher-http-go added at 8d4ffe1576