ollie/toolsrv/shell_validate.go

87 lines
2.4 KiB
Go

package toolsrv
import (
"fmt"
"regexp"
"strings"
"time"
)
// universalPatterns apply to all code.
var universalPatterns = []*regexp.Regexp{
regexp.MustCompile(`\bmkfs\b`),
regexp.MustCompile(`\bdd\b.*\bif=/dev/`),
regexp.MustCompile(`\b(sudo|su)\s`),
regexp.MustCompile(`/etc/(shadow|sudoers)`),
}
// bashPatterns apply to bash (flag syntax, redirects, shell-specific constructs).
var bashPatterns = []*regexp.Regexp{
regexp.MustCompile(`rm\s+(-[a-z]*r[a-z]*\s+)*-[a-z]*f[a-z]*\s*/(home|var|usr|etc|boot|root|bin|sbin|lib|opt|srv)?`),
regexp.MustCompile(`rm\s+(-[a-z]*f[a-z]*\s+)*-[a-z]*r[a-z]*\s*/(home|var|usr|etc|boot|root|bin|sbin|lib|opt|srv)?`),
regexp.MustCompile(`rm\s+.*--recursive.*--force`),
regexp.MustCompile(`rm\s+.*--force.*--recursive`),
regexp.MustCompile(`rm\s+(-[a-z]*r[a-z]*\s+)*-[a-z]*f[a-z]*\s+\.\.(/|\s|$)`),
regexp.MustCompile(`rm\s+(-[a-z]*r[a-z]*\s+)*-[a-z]*f[a-z]*\s+~`),
regexp.MustCompile(`rm\s+(-[a-z]*r[a-z]*\s+)*-[a-z]*f[a-z]*\s+\*`),
regexp.MustCompile(`:\s*\(\s*\)\s*\{\s*:\s*\|\s*:\s*&`), // fork bomb
regexp.MustCompile(`>\s*/dev/sd`),
regexp.MustCompile(`\beval\s+".*\$`),
}
var languagePatterns = map[string][]*regexp.Regexp{
"bash": bashPatterns,
"": bashPatterns,
}
var whitespacePattern = regexp.MustCompile(`\s+`)
// ValidateCode checks code against dangerous patterns.
func (s *Server) ValidateCode(code, language string) error {
if err := s.checkRateLimit(); err != nil {
return err
}
normalized := strings.ToLower(code)
normalized = whitespacePattern.ReplaceAllString(normalized, " ")
patterns := append(universalPatterns, languagePatterns[language]...)
for _, pattern := range patterns {
if pattern.MatchString(normalized) {
s.recordValidationFailure()
return fmt.Errorf("dangerous pattern detected")
}
}
return nil
}
func (s *Server) checkRateLimit() error {
s.rateLimitMu.Lock()
defer s.rateLimitMu.Unlock()
now := time.Now()
if now.Before(s.blockedUntil) {
remaining := s.blockedUntil.Sub(now).Round(time.Second)
return &RateLimitedError{Remaining: remaining}
}
return nil
}
func (s *Server) recordValidationFailure() {
s.rateLimitMu.Lock()
defer s.rateLimitMu.Unlock()
now := time.Now()
if now.Sub(s.lastFailure) > failureWindow {
s.validationFailures = 0
}
s.validationFailures++
s.lastFailure = now
if s.validationFailures >= maxFailures {
s.blockedUntil = now.Add(blockDuration)
s.validationFailures = 0
}
}