// Copyright (C) 2019 Marius Schellenberger package main import ( "bytes" "encoding/json" "errors" "fmt" "io/ioutil" "log" "net" "net/http" "net/url" "reflect" "regexp" "strings" "time" //corev1 "k8s.io/api/core/v1" metav1 "k8s.io/apimachinery/pkg/apis/meta/v1" "k8s.io/client-go/kubernetes" typedcorev1 "k8s.io/client-go/kubernetes/typed/core/v1" "k8s.io/client-go/rest" ) var debug = false const ( prAction = "X-Interceptor-Pr-Action" pushAction = "X-Interceptor-Push-Action" commentAction = "X-Interceptor-Comment-Action" labelHeader = "X-Interceptor-Label" removeLabelsHeader = "X-Interceptor-Remove-Labels" contentType = "Content-Type" jsonType = "application/json" githubEvent = "X-GitHub-Event" giteaEvent = "X-Gitea-Event" issueHeader = "Issue" commitHeader = "Commit" secretName = "git-auth" ) func init() { log.SetFlags(log.Flags() | log.Lshortfile) } func (h *Handler) getSecret(host string) (username, password string, err error) { sec, err := h.sec.Get(secretName, metav1.GetOptions{}) if err != nil { return } s, ok := sec.Data[host] if !ok { err = fmt.Errorf("no key found for host %s in secret %s/%s", host, sec.Namespace, sec.Name) return } a := strings.SplitN(string(s), ":", 2) if len(a) < 2 { err = fmt.Errorf("missing credentials for host %s in secret %s/%s", host, sec.Namespace, sec.Name) return } username = a[0] password = a[1] return } func (h *Handler) getPullRequest(repo Repository, api, id string) (*PullRequest, error) { u, err := url.Parse(repo.CloneURL) if err != nil { return nil, err } u.Path = strings.Join([]string{api, "repos", repo.Owner.Login, repo.Name, "pulls", id}, "/") req, err := http.NewRequest("GET", u.String(), nil) if err != nil { return nil, err } req.Header.Set(contentType, jsonType) user, pass, err := h.getSecret(u.Host) if err != nil { return nil, err } req.SetBasicAuth(user, pass) resp, err := h.cl.Do(req) if err != nil { return nil, err } if resp.StatusCode != http.StatusOK { return nil, errors.New("error getting pull request") } pr := new(PullRequest) _, err = readBody(resp.Body, pr) if err != nil { return nil, err } return pr, nil } func (h *Handler) addLabel(repo Repository, api, label, color string) (labelID int64, err error) { labelID = -1 u, err := url.Parse(repo.CloneURL) if err != nil { err = fmt.Errorf("error parsing repo URL: %s", err) return } repoLabels := strings.Join([]string{api, "repos", repo.Owner.Login, repo.Name, "labels"}, "/") u.Path = repoLabels req, err := http.NewRequest("GET", u.String(), nil) if err != nil { return } user, pass, err := h.getSecret(u.Host) if err != nil { return } req.SetBasicAuth(user, pass) req.Header.Set(contentType, jsonType) resp, err := h.cl.Do(req) if err != nil { err = fmt.Errorf("error getting labels of %s: %s", u, err) return } if resp.StatusCode != http.StatusOK { err = fmt.Errorf("error getting labels of %s", u) return } var labels []Label _, err = readBody(resp.Body, &labels) if err != nil { return } for _, l := range labels { if l.Name == label { labelID = l.ID break } } if labelID >= 0 { return } buf := new(bytes.Buffer) l := &Label{Name: label, Color: color} err = json.NewEncoder(buf).Encode(l) if err != nil { return } req, err = http.NewRequest("POST", u.String(), buf) if err != nil { return } req.Header.Set(contentType, jsonType) req.SetBasicAuth(user, pass) resp, err = h.cl.Do(req) if err != nil { err = fmt.Errorf("error creating label %s for %s: %s", label, u, err) return } if resp.StatusCode != http.StatusCreated { err = fmt.Errorf("error creating label %s for %s", label, u) } var createdLabel Label _, err = readBody(resp.Body, &createdLabel) labelID = createdLabel.ID return } func (h *Handler) setLabel(repo Repository, api, id, label, color string, remove []string) (err error) { labelID, err := h.addLabel(repo, api, label, color) if err != nil { return } u, err := url.Parse(repo.CloneURL) if err != nil { return } u.Path = strings.Join([]string{api, "repos", repo.Owner.Login, repo.Name, "issues", id, "labels"}, "/") req, err := http.NewRequest("GET", u.String(), nil) if err != nil { return } user, pass, err := h.getSecret(u.Host) if err != nil { return } req.SetBasicAuth(user, pass) req.Header.Set(contentType, jsonType) resp, err := h.cl.Do(req) if err != nil { return fmt.Errorf("error getting issue labels of %s: %s", u, err) } if resp.StatusCode != http.StatusOK { return fmt.Errorf("error getting issue labels of %s", u) } var labels []Label _, err = readBody(resp.Body, &labels) if err != nil { return } pl := &PostLabels{[]int64{labelID}} var rm, found bool for _, l := range labels { if contains(remove, l.Name) { rm = true } else { if l.Name == label { found = true } pl.Labels = append(pl.Labels, l.ID) } } if !rm && found { return } buf := new(bytes.Buffer) err = json.NewEncoder(buf).Encode(pl) if err != nil { return } req, err = http.NewRequest("PUT", u.String(), buf) if err != nil { return } req.SetBasicAuth(user, pass) req.Header.Set(contentType, jsonType) resp, err = h.cl.Do(req) if err != nil { return fmt.Errorf("error setting issue label %s for %s: %s", label, u, err) } if resp.StatusCode != http.StatusOK { err = fmt.Errorf("error setting issue label %s for %s", label, u) } return } type Handler struct { sec typedcorev1.SecretInterface cl *http.Client } func (h *Handler) ServeHTTP(w http.ResponseWriter, r *http.Request) { event := r.Header.Get(githubEvent) switch event { case "pull_request": h.prHandler(w, r) case "push": h.pushHandler(w, r) case "issue_comment": h.commentHandler(w, r) default: h.dumpHandler(w, r) } status := int(reflect.Indirect(reflect.ValueOf(w)).FieldByName("status").Int()) log.Printf("%s %s %s event:%s status:%d", r.RemoteAddr, r.Method, r.URL, event, status) } func (h *Handler) prHandler(w http.ResponseWriter, r *http.Request) { var prb PullRequestBody body, err := readBody(r.Body, &prb) if err != nil { logError(w, err.Error()) return } if contains(r.Header[prAction], prb.Action) { if l := r.Header.Get(labelHeader); l != "" { label, _ := getLabel(l) var labels []string for _, l := range prb.PullRequest.Labels { labels = append(labels, l.Name) } if !contains(labels, label) { w.WriteHeader(http.StatusBadRequest) return } } h := w.Header() h.Set(issueHeader, itoa(prb.Number)) h.Set(commitHeader, prb.PullRequest.Head.SHA) w.Write(body) return } w.WriteHeader(http.StatusBadRequest) } func (h *Handler) pushHandler(w http.ResponseWriter, r *http.Request) { var pb PushBody body, err := readBody(r.Body, &pb) if err != nil { logError(w, err.Error()) return } push := r.Header.Get(pushAction) //if contains(r.Header[pushAction], pb.Ref) { if push != "" && push == pb.Ref { h := w.Header() h.Set(commitHeader, pb.After) w.Write(body) return } w.WriteHeader(http.StatusBadRequest) } func (h *Handler) commentHandler(w http.ResponseWriter, r *http.Request) { var cb CommentBody body, err := readBody(r.Body, &cb) if err != nil { logError(w, err.Error()) return } //if contains(r.Header[commentAction], cb.Comment.Body) { comment := r.Header.Get(commentAction) if comment == "" { w.WriteHeader(http.StatusBadRequest) return } m, err := regexp.MatchString("(?m)^"+comment+"$", cb.Comment.Body) if err != nil { logError(w, err.Error()) return } if m { issue := itoa(cb.Issue.Number) api := detectAPI(r) if api == "" { logError(w, "unable to detect git provider API") return } if l := r.Header.Get(labelHeader); l != "" { label, color := getLabel(l) var labels []string for _, l := range cb.Issue.Labels { labels = append(labels, l.Name) } if !contains(labels, label) { err := h.setLabel(cb.Repository, api, issue, label, color, r.Header[removeLabelsHeader]) if err != nil { logError(w, err.Error()) return } } } pr, err := h.getPullRequest(cb.Repository, api, issue) if err != nil { logError(w, err.Error()) return } h := w.Header() h.Set(issueHeader, issue) h.Set(commitHeader, pr.Head.SHA) w.Write(body) return } w.WriteHeader(http.StatusBadRequest) } func (h *Handler) dumpHandler(w http.ResponseWriter, r *http.Request) { defer r.Body.Close() if debug { body, err := ioutil.ReadAll(r.Body) if err != nil { logError(w, err.Error()) return } log.Printf("debug headers: %#v", r.Header) log.Printf("debug: %s", string(body)) } w.WriteHeader(http.StatusBadRequest) } func main() { cfg, err := rest.InClusterConfig() if err != nil { log.Fatalf("error configuring kube client: %s\n", err) } cl, err := kubernetes.NewForConfig(cfg) if err != nil { log.Fatalf("error creating new kube client: %s\n", err) } ns, err := getNS() if err != nil { log.Fatalf("error getting current namespace %s\n", err) } h := &Handler{ sec: cl.CoreV1().Secrets(ns), cl: &http.Client{Transport: &http.Transport{ DialContext: (&net.Dialer{ Timeout: 30 * time.Second, }).DialContext, }}, } http.Handle("/", h) s := &http.Server{ ReadTimeout: 5 * time.Second, WriteTimeout: 10 * time.Second, IdleTimeout: 120 * time.Second, Addr: ":8080", } log.Println("started webhook listener") log.Fatal(s.ListenAndServe()) }