87 lines
2.4 KiB
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
|
|
}
|
|
}
|