Files
ThothII/tools/tht/internal/preflight/network.go
T

217 lines
7.3 KiB
Go

package preflight
import (
"context"
"crypto/tls"
"crypto/x509"
"errors"
"io"
"net"
"net/http"
"net/url"
"strconv"
"strings"
"time"
"github.com/aritmolab/thothii/tools/tht/internal/config"
"github.com/aritmolab/thothii/tools/tht/internal/safeio"
"golang.org/x/crypto/ssh"
"golang.org/x/crypto/ssh/knownhosts"
)
func CheckExternal(ctx context.Context, installation config.Installation) Report {
r := NewReport()
value := func(name string) string { result, _ := installation.EnvironmentValue(name); return result }
err := checkGit(ctx, installation, value)
outcome := "passed"
if err != nil {
outcome = "error"
}
r.Add("workspace-remote", outcome, "workspaceRepository", "Require authenticated read access to the configured Git remote and branch using the prepared trust/credential files.")
for name, provider := range installation.ModelCatalog.Providers {
if provider.Endpoint == nil {
r.Add("provider-"+name, "warning", "modelCatalog.providers", "Built-in provider endpoint resolution belongs to the bundled Pi SDK. The required Pi runtime smoke check verifies model availability and credentials; preflight makes no billable generation requests.")
continue
}
parsed, err := url.Parse(provider.Endpoint.BaseURL)
if err == nil {
err = probeOrigin(ctx, parsed)
}
outcome := "passed"
if err != nil {
outcome = "error"
}
r.Add("provider-"+name, outcome, "modelCatalog.providers."+name+".endpoint", "Require DNS/TCP/TLS reachability of the configured provider origin. Credential/model eligibility still requires the bundled Pi runtime smoke check; no generation request is sent here.")
}
return r
}
func probeOrigin(ctx context.Context, target *url.URL) error {
bound, cancel := context.WithTimeout(ctx, 5*time.Second)
defer cancel()
port := target.Port()
if port == "" {
port = "443"
if target.Scheme == "http" {
port = "80"
}
}
address := net.JoinHostPort(target.Hostname(), port)
var connection net.Conn
var err error
if target.Scheme == "https" {
dialer := tls.Dialer{NetDialer: &net.Dialer{Timeout: 5 * time.Second}, Config: &tls.Config{MinVersion: tls.VersionTLS12, ServerName: target.Hostname()}}
connection, err = dialer.DialContext(bound, "tcp", address)
} else {
connection, err = (&net.Dialer{Timeout: 5 * time.Second}).DialContext(bound, "tcp", address)
}
if err == nil {
connection.Close()
}
return err
}
func checkGit(ctx context.Context, installation config.Installation, value func(string) string) error {
remote := installation.WorkspaceRepository.Remote
if installation.WorkspaceRepository.Access == "ssh" && strings.HasPrefix(remote, "git@") {
host, path, found := strings.Cut(strings.TrimPrefix(remote, "git@"), ":")
if !found {
return errors.New("invalid Git remote")
}
return checkSSHGit(ctx, &url.URL{Scheme: "ssh", User: url.User("git"), Host: host, Path: path}, installation.WorkspaceRepository.Branch, value)
}
u, err := url.Parse(remote)
if err != nil {
return errors.New("remote unavailable")
}
if installation.WorkspaceRepository.Access == "ssh" {
return checkSSHGit(ctx, u, installation.WorkspaceRepository.Branch, value)
}
ca, err := safeio.ReadCanonicalPrivateRegular(value("THT_WORKSPACE_GIT_CA_FILE"), 64<<10)
if err != nil {
return err
}
pool, err := x509.SystemCertPool()
if err != nil {
pool = x509.NewCertPool()
}
if !pool.AppendCertsFromPEM(ca) {
return errors.New("invalid Git trust")
}
credentials, err := safeio.ReadCanonicalPrivateRegular(value("THT_WORKSPACE_GIT_CREDENTIALS_FILE"), 64<<10)
if err != nil {
return err
}
var user, password string
for _, line := range strings.Split(string(credentials), "\n") {
if strings.TrimSpace(line) == "" {
continue
}
credential, err := url.Parse(strings.TrimSpace(line))
if err != nil || credential.User == nil {
return errors.New("invalid Git credentials")
}
if credential.Scheme == u.Scheme && credential.Host == u.Host && (credential.Path == "" || credential.Path == u.Path) {
user = credential.User.Username()
password, _ = credential.User.Password()
}
}
u.Path = strings.TrimSuffix(u.Path, "/") + "/info/refs"
u.RawQuery = "service=git-upload-pack"
bound, cancel := context.WithTimeout(ctx, 5*time.Second)
defer cancel()
request, err := http.NewRequestWithContext(bound, http.MethodGet, u.String(), nil)
if err != nil {
return err
}
if user != "" {
request.SetBasicAuth(user, password)
}
transport := &http.Transport{TLSClientConfig: &tls.Config{RootCAs: pool, MinVersion: tls.VersionTLS12}, Proxy: http.ProxyFromEnvironment}
defer transport.CloseIdleConnections()
client := http.Client{Transport: transport, Timeout: 5 * time.Second, CheckRedirect: func(*http.Request, []*http.Request) error { return http.ErrUseLastResponse }}
response, err := client.Do(request)
if err != nil {
return err
}
defer response.Body.Close()
if response.StatusCode != http.StatusOK {
return errors.New("Git remote refused")
}
data, err := io.ReadAll(io.LimitReader(response.Body, (1<<20)+1))
if err != nil || len(data) > 1<<20 || !advertisesBranch(data, installation.WorkspaceRepository.Branch) {
return errors.New("Git branch unavailable")
}
return nil
}
func advertisesBranch(data []byte, branch string) bool {
for _, line := range strings.Split(string(data), "\n") {
line = strings.SplitN(line, "\x00", 2)[0]
if strings.HasSuffix(line, " refs/heads/"+branch) {
return true
}
}
return false
}
func checkSSHGit(ctx context.Context, target *url.URL, branch string, value func(string) string) error {
// Git paths and branch names are already validated by the canonical installation loader.
if target.Scheme != "ssh" || target.User == nil || strings.ContainsAny(target.Path, "'\r\n\x00") {
return errors.New("use canonical ssh:// remote")
}
key, err := safeio.ReadCanonicalPrivateRegular(value("THT_WORKSPACE_GIT_SSH_KEY_FILE"), 64<<10)
if err != nil {
return err
}
signer, err := ssh.ParsePrivateKey(key)
if err != nil {
return err
}
hostKey, err := knownhosts.New(value("THT_WORKSPACE_GIT_KNOWN_HOSTS_FILE"))
if err != nil {
return err
}
port := target.Port()
if port == "" {
port = "22"
}
if _, err := strconv.Atoi(port); err != nil {
return err
}
address := net.JoinHostPort(target.Hostname(), port)
bound, cancel := context.WithTimeout(ctx, 5*time.Second)
defer cancel()
connection, err := (&net.Dialer{Timeout: 5 * time.Second}).DialContext(bound, "tcp", address)
if err != nil {
return err
}
defer connection.Close()
deadline, _ := bound.Deadline()
_ = connection.SetDeadline(deadline)
clientConnection, channels, requests, err := ssh.NewClientConn(connection, address, &ssh.ClientConfig{User: target.User.Username(), Auth: []ssh.AuthMethod{ssh.PublicKeys(signer)}, HostKeyCallback: hostKey, Timeout: 5 * time.Second})
if err != nil {
return err
}
client := ssh.NewClient(clientConnection, channels, requests)
defer client.Close()
session, err := client.NewSession()
if err != nil {
return err
}
defer session.Close()
stdout, err := session.StdoutPipe()
if err != nil {
return err
}
stdin, err := session.StdinPipe()
if err != nil {
return err
}
if err = session.Start("git-upload-pack --advertise-refs '" + target.Path + "'"); err != nil {
return err
}
_ = stdin.Close()
data, err := io.ReadAll(io.LimitReader(stdout, (1<<20)+1))
if err != nil || len(data) > 1<<20 || !advertisesBranch(data, branch) {
return errors.New("Git branch unavailable")
}
return nil
}