491 lines
15 KiB
Go
491 lines
15 KiB
Go
package main
|
|
|
|
import (
|
|
"crypto/sha256"
|
|
"encoding/hex"
|
|
"encoding/json"
|
|
"errors"
|
|
"fmt"
|
|
"io"
|
|
"os"
|
|
"path/filepath"
|
|
"slices"
|
|
"strings"
|
|
"time"
|
|
)
|
|
|
|
const aiStoreVersion = 1
|
|
|
|
type AIStore struct {
|
|
dir string
|
|
loadErr error
|
|
}
|
|
|
|
type aiStoredState struct {
|
|
Version int `json:"version"`
|
|
Owner string `json:"owner"`
|
|
Repository string `json:"repository"`
|
|
Number int `json:"number"`
|
|
Threads []ReviewThread `json:"threads"`
|
|
Annotations map[string][]ReviewComment `json:"annotations"`
|
|
}
|
|
|
|
func NewAIStore(dir string) *AIStore { return &AIStore{dir: dir} }
|
|
|
|
func (s *AIStore) path(pr PRDetails) string {
|
|
sum := sha256.Sum256([]byte(fmt.Sprintf("%s/%s#%d", pr.Owner, pr.Repository, pr.Number)))
|
|
return filepath.Join(s.dir, hex.EncodeToString(sum[:16])+".json")
|
|
}
|
|
|
|
func (s *AIStore) Load(pr PRDetails) (*aiStoredState, error) {
|
|
state := &aiStoredState{
|
|
Version: aiStoreVersion, Owner: pr.Owner, Repository: pr.Repository, Number: pr.Number,
|
|
Annotations: make(map[string][]ReviewComment),
|
|
}
|
|
if s == nil || s.dir == "" {
|
|
return state, nil
|
|
}
|
|
file, err := os.Open(s.path(pr))
|
|
if errors.Is(err, os.ErrNotExist) {
|
|
return state, nil
|
|
}
|
|
if err != nil {
|
|
s.loadErr = err
|
|
return nil, fmt.Errorf("read local AI state: %w", err)
|
|
}
|
|
defer file.Close()
|
|
data, err := io.ReadAll(io.LimitReader(file, (16<<20)+1))
|
|
if err != nil {
|
|
return nil, fmt.Errorf("read local AI state: %w", err)
|
|
}
|
|
if len(data) > 16<<20 {
|
|
return nil, errors.New("read local AI state: file exceeds 16 MiB safety limit")
|
|
}
|
|
if err := json.Unmarshal(data, state); err != nil || state.Version != aiStoreVersion ||
|
|
state.Owner != pr.Owner || state.Repository != pr.Repository || state.Number != pr.Number {
|
|
if err == nil {
|
|
err = errors.New("incompatible or mismatched state")
|
|
}
|
|
s.loadErr = err
|
|
return nil, fmt.Errorf("read local AI state: %w", err)
|
|
}
|
|
if state.Annotations == nil {
|
|
state.Annotations = make(map[string][]ReviewComment)
|
|
}
|
|
return state, nil
|
|
}
|
|
|
|
func (s *AIStore) Save(pr PRDetails, state *aiStoredState) error {
|
|
if s == nil || s.dir == "" {
|
|
return errors.New("local AI store is unavailable")
|
|
}
|
|
if err := os.MkdirAll(s.dir, 0o700); err != nil {
|
|
return fmt.Errorf("create local AI store: %w", err)
|
|
}
|
|
if err := os.Chmod(s.dir, 0o700); err != nil {
|
|
return fmt.Errorf("protect local AI store: %w", err)
|
|
}
|
|
path := s.path(pr)
|
|
next, err := json.MarshalIndent(state, "", " ")
|
|
if err != nil {
|
|
return err
|
|
}
|
|
if current, err := os.ReadFile(path); err == nil && string(current) == string(next) {
|
|
return nil
|
|
}
|
|
return atomicWriteJSON(path, state, 0o600)
|
|
}
|
|
|
|
func (s *aiStoredState) Merge(pr PRDetails) PRDetails {
|
|
result := pr
|
|
result.Threads = make([]ReviewThread, len(pr.Threads), len(pr.Threads)+len(s.Threads))
|
|
copy(result.Threads, pr.Threads)
|
|
for index := range result.Threads {
|
|
result.Threads[index].Comments = slices.Clone(result.Threads[index].Comments)
|
|
}
|
|
for i := range s.Threads {
|
|
thread := s.Threads[i]
|
|
thread.Comments = slices.Clone(thread.Comments)
|
|
thread.IsOutdated = thread.HeadOID != "" && thread.HeadOID != pr.HeadOID
|
|
result.Threads = append(result.Threads, thread)
|
|
}
|
|
for i := range result.Threads {
|
|
if comments := s.Annotations[result.Threads[i].ID]; len(comments) > 0 {
|
|
comments = slices.Clone(comments)
|
|
for index := range comments {
|
|
if comments[index].Origin == reviewOriginLocalAIUser && pr.ViewerLogin != "" {
|
|
comments[index].Author = pr.ViewerLogin
|
|
}
|
|
}
|
|
result.Threads[i].Comments = append(result.Threads[i].Comments, comments...)
|
|
slices.SortStableFunc(result.Threads[i].Comments, func(left, right ReviewComment) int {
|
|
return left.CreatedAt.Compare(right.CreatedAt)
|
|
})
|
|
}
|
|
}
|
|
return result
|
|
}
|
|
|
|
func (s *aiStoredState) Apply(
|
|
pr PRDetails, output aiOutput, provider, model, targetThread, message string,
|
|
validLines map[string]map[int]bool,
|
|
validDeleted map[string]map[int]bool,
|
|
diffText map[string]string,
|
|
) (int, int) {
|
|
existing := make(map[string]bool)
|
|
localByFingerprint := make(map[string]int, len(s.Threads))
|
|
for index, thread := range s.Threads {
|
|
if thread.Fingerprint != "" {
|
|
localByFingerprint[thread.Fingerprint] = index
|
|
}
|
|
}
|
|
type priorFinding struct {
|
|
path string
|
|
start, end int
|
|
text string
|
|
}
|
|
var prior []priorFinding
|
|
for _, thread := range append(slices.Clone(pr.Threads), s.Threads...) {
|
|
existing[thread.Fingerprint] = thread.Fingerprint != ""
|
|
var combined strings.Builder
|
|
for _, comment := range thread.Comments {
|
|
existing[aiFingerprint(thread.Path, thread.StartLine, thread.Line, "", comment.Body)] = true
|
|
existing[aiFingerprint(thread.ID, 0, 0, "", comment.Body)] = true
|
|
if comment.Origin == reviewOriginLocalAIUser {
|
|
existing[aiFingerprint(thread.ID, 0, 0, "user", comment.Body)] = true
|
|
}
|
|
combined.WriteString(" ")
|
|
combined.WriteString(comment.Body)
|
|
}
|
|
prior = append(prior, priorFinding{
|
|
path: thread.Path, start: thread.StartLine, end: thread.Line,
|
|
text: combined.String(),
|
|
})
|
|
}
|
|
findings := 0
|
|
if targetThread != "" {
|
|
output.Findings = nil
|
|
}
|
|
for _, finding := range output.Findings {
|
|
finding.Path = safeAIText(finding.Path)
|
|
finding.Title = safeAIText(finding.Title)
|
|
finding.Body = safeAIText(finding.Body)
|
|
if finding.EndLine < finding.StartLine {
|
|
finding.EndLine = finding.StartLine
|
|
}
|
|
if finding.EndLine-finding.StartLine > 500 || finding.EndLine > 10_000_000 {
|
|
continue
|
|
}
|
|
changed, validPath := validLines[finding.Path]
|
|
if finding.Side == "LEFT" {
|
|
changed, validPath = validDeleted[finding.Path]
|
|
}
|
|
ok := false
|
|
for line := finding.StartLine; validPath && line <= finding.EndLine; line++ {
|
|
ok = ok || changed[line]
|
|
}
|
|
if !validPath || !ok {
|
|
continue
|
|
}
|
|
if finding.Path == "" || (finding.Side != "LEFT" && finding.Side != "RIGHT") ||
|
|
finding.StartLine < 1 || finding.Body == "" || len(finding.Body) > 16_000 {
|
|
continue
|
|
}
|
|
fingerprint := aiFingerprint(finding.Path, finding.StartLine, finding.EndLine, finding.Title, finding.Body)
|
|
suggestion := validatedAISuggestion(finding, changed)
|
|
if existing[fingerprint] {
|
|
if index, ok := localByFingerprint[fingerprint]; ok && suggestion != "" &&
|
|
len(s.Threads[index].Comments) > 0 &&
|
|
len(parseCommentBody(s.Threads[index].Comments[0].Body).Suggestions) == 0 {
|
|
s.Threads[index].Comments[0].Body +=
|
|
"\n\n```suggestion\n" + suggestion + "\n```"
|
|
}
|
|
continue
|
|
}
|
|
nearDuplicate := false
|
|
for _, item := range prior {
|
|
if item.path == finding.Path &&
|
|
rangesNear(item.start, item.end, finding.StartLine, finding.EndLine) &&
|
|
aiTextSimilarity(item.text, finding.Title+" "+finding.Body) >= 0.68 {
|
|
nearDuplicate = true
|
|
break
|
|
}
|
|
}
|
|
if nearDuplicate {
|
|
continue
|
|
}
|
|
existing[fingerprint] = true
|
|
id := "local-ai-" + fingerprint
|
|
author := model
|
|
body := "**" + finding.Severity + ": " + finding.Title + "**\n\n" + finding.Body
|
|
if suggestion != "" {
|
|
body += "\n\n```suggestion\n" + suggestion + "\n```"
|
|
}
|
|
s.Threads = append(s.Threads, ReviewThread{
|
|
ID: id, Path: finding.Path, StartLine: finding.StartLine, Line: finding.EndLine,
|
|
DiffSide: finding.Side, Origin: reviewOriginLocalAI, Provider: provider, Model: model,
|
|
HeadOID: pr.HeadOID, Fingerprint: fingerprint, ViewerCanResolve: true,
|
|
ViewerCanUnresolve: true, ViewerCanReply: true,
|
|
Comments: []ReviewComment{{
|
|
ID: id + "-0", Author: author, Body: body,
|
|
CreatedAt: time.Now(), Origin: reviewOriginLocalAI, Provider: provider, Model: model,
|
|
DiffHunk: boundedDiffHunk(diffText[finding.Path], finding.StartLine, finding.Side),
|
|
}},
|
|
})
|
|
findings++
|
|
prior = append(prior, priorFinding{
|
|
path: finding.Path, start: finding.StartLine, end: finding.EndLine,
|
|
text: finding.Title + " " + finding.Body,
|
|
})
|
|
}
|
|
comments := 0
|
|
validThreads := make(map[string]bool)
|
|
for _, thread := range pr.Threads {
|
|
validThreads[thread.ID] = !thread.IsResolved
|
|
}
|
|
if targetThread != "" {
|
|
validThreads = map[string]bool{targetThread: true}
|
|
}
|
|
message = safeAIText(message)
|
|
if targetThread != "" && strings.TrimSpace(message) != "" {
|
|
fingerprint := aiFingerprint(targetThread, 0, 0, "user", message)
|
|
if !existing[fingerprint] {
|
|
existing[fingerprint] = true
|
|
s.Annotations[targetThread] = append(s.Annotations[targetThread], ReviewComment{
|
|
ID: "local-ai-user-" + fingerprint,
|
|
Author: firstNonEmpty(pr.ViewerLogin, "you"), Body: message,
|
|
CreatedAt: time.Now(), Origin: reviewOriginLocalAIUser,
|
|
})
|
|
}
|
|
}
|
|
for _, annotation := range output.ThreadComments {
|
|
if !validThreads[annotation.ThreadID] {
|
|
continue
|
|
}
|
|
body := safeAIText(annotation.Body)
|
|
if body == "" || len(body) > 16_000 {
|
|
continue
|
|
}
|
|
fingerprint := aiFingerprint(annotation.ThreadID, 0, 0, "", body)
|
|
if existing[fingerprint] {
|
|
continue
|
|
}
|
|
existing[fingerprint] = true
|
|
s.Annotations[annotation.ThreadID] = append(s.Annotations[annotation.ThreadID], ReviewComment{
|
|
ID: "local-ai-comment-" + fingerprint, Author: model, Body: body,
|
|
CreatedAt: time.Now(), Origin: reviewOriginLocalAI, Provider: provider, Model: model,
|
|
})
|
|
comments++
|
|
}
|
|
sortAIThreads(s.Threads)
|
|
return findings, comments
|
|
}
|
|
|
|
func validatedAISuggestion(finding aiFinding, changed map[int]bool) string {
|
|
suggestion := strings.Trim(sanitizeAIControls(finding.Suggestion), "\r\n")
|
|
if suggestion == "" || finding.Side != "RIGHT" ||
|
|
finding.EndLine-finding.StartLine+1 > 12 ||
|
|
len(suggestion) > 8_000 || strings.Contains(suggestion, "```") {
|
|
return ""
|
|
}
|
|
for line := finding.StartLine; line <= finding.EndLine; line++ {
|
|
if !changed[line] {
|
|
return ""
|
|
}
|
|
}
|
|
replacementLines := strings.Split(suggestion, "\n")
|
|
if len(replacementLines) > 12 {
|
|
return ""
|
|
}
|
|
for _, line := range replacementLines {
|
|
if len(line) > 2_000 {
|
|
return ""
|
|
}
|
|
}
|
|
return suggestion
|
|
}
|
|
|
|
func rangesNear(leftStart, leftEnd, rightStart, rightEnd int) bool {
|
|
if leftStart == 0 {
|
|
leftStart = leftEnd
|
|
}
|
|
if rightStart == 0 {
|
|
rightStart = rightEnd
|
|
}
|
|
return leftStart <= rightEnd+3 && rightStart <= leftEnd+3
|
|
}
|
|
|
|
func aiTextSimilarity(left, right string) float64 {
|
|
tokenize := func(value string) map[string]bool {
|
|
tokens := make(map[string]bool)
|
|
stop := map[string]bool{
|
|
"and": true, "are": true, "can": true, "for": true, "from": true,
|
|
"that": true, "the": true, "this": true, "when": true, "with": true,
|
|
}
|
|
for _, token := range strings.FieldsFunc(strings.ToLower(value), func(r rune) bool {
|
|
return !(r >= 'a' && r <= 'z' || r >= '0' && r <= '9' || r == '_')
|
|
}) {
|
|
if len(token) >= 3 && !stop[token] {
|
|
tokens[token] = true
|
|
}
|
|
}
|
|
return tokens
|
|
}
|
|
a, b := tokenize(left), tokenize(right)
|
|
if len(a) < 3 || len(b) < 3 {
|
|
return 0
|
|
}
|
|
intersection := 0
|
|
for token := range a {
|
|
if b[token] {
|
|
intersection++
|
|
}
|
|
}
|
|
return float64(intersection) / float64(min(len(a), len(b)))
|
|
}
|
|
|
|
func boundedDiffHunk(diff string, targetLine int, side string) string {
|
|
lines := strings.Split(diff, "\n")
|
|
for start := 0; start < len(lines); start++ {
|
|
if !strings.HasPrefix(lines[start], "@@ ") {
|
|
continue
|
|
}
|
|
fields := strings.Fields(lines[start])
|
|
if len(fields) < 3 {
|
|
continue
|
|
}
|
|
value := strings.TrimPrefix(strings.SplitN(fields[2], ",", 2)[0], "+")
|
|
newLine := 0
|
|
fmt.Sscanf(value, "%d", &newLine)
|
|
oldValue := strings.TrimPrefix(strings.SplitN(fields[1], ",", 2)[0], "-")
|
|
oldLine := 0
|
|
fmt.Sscanf(oldValue, "%d", &oldLine)
|
|
end := len(lines)
|
|
for index := start + 1; index < len(lines); index++ {
|
|
if strings.HasPrefix(lines[index], "@@ ") {
|
|
end = index
|
|
break
|
|
}
|
|
}
|
|
for index := start + 1; index < end; index++ {
|
|
line := lines[index]
|
|
if line == "" {
|
|
oldLine++
|
|
newLine++
|
|
continue
|
|
}
|
|
current := newLine
|
|
if side == "LEFT" {
|
|
current = oldLine
|
|
}
|
|
if line[0] != '+' {
|
|
oldLine++
|
|
}
|
|
if line[0] != '-' {
|
|
newLine++
|
|
}
|
|
targetSideLine := side == "LEFT" && line[0] != '+' ||
|
|
side == "RIGHT" && line[0] != '-'
|
|
if current == targetLine && targetSideLine {
|
|
hunk := strings.Join(lines[start:end], "\n")
|
|
if len(hunk) > 24_000 {
|
|
hunk = compactDiffHunk(lines, start, end, index)
|
|
}
|
|
return sanitizeAIControls(hunk)
|
|
}
|
|
}
|
|
start = end - 1
|
|
}
|
|
return ""
|
|
}
|
|
|
|
func compactDiffHunk(lines []string, hunkStart, hunkEnd, target int) string {
|
|
fields := strings.Fields(lines[hunkStart])
|
|
if len(fields) < 3 {
|
|
return ""
|
|
}
|
|
oldStart, newStart := 0, 0
|
|
fmt.Sscanf(strings.TrimPrefix(strings.SplitN(fields[1], ",", 2)[0], "-"), "%d", &oldStart)
|
|
fmt.Sscanf(strings.TrimPrefix(strings.SplitN(fields[2], ",", 2)[0], "+"), "%d", &newStart)
|
|
windowStart := max(hunkStart+1, target-20)
|
|
windowEnd := min(hunkEnd, target+21)
|
|
for index := hunkStart + 1; index < windowStart; index++ {
|
|
line := lines[index]
|
|
if line == "" || line[0] != '+' {
|
|
oldStart++
|
|
}
|
|
if line == "" || line[0] != '-' {
|
|
newStart++
|
|
}
|
|
}
|
|
oldCount, newCount := 0, 0
|
|
body := make([]string, 0, windowEnd-windowStart)
|
|
for _, line := range lines[windowStart:windowEnd] {
|
|
if line == "" || line[0] != '+' {
|
|
oldCount++
|
|
}
|
|
if line == "" || line[0] != '-' {
|
|
newCount++
|
|
}
|
|
if len(line) > 500 {
|
|
prefix := ""
|
|
if line != "" {
|
|
prefix = line[:1]
|
|
line = line[1:]
|
|
}
|
|
line = prefix + truncateAIInline(line, 480)
|
|
}
|
|
body = append(body, line)
|
|
}
|
|
header := fmt.Sprintf("@@ -%d,%d +%d,%d @@ local snapshot", oldStart, oldCount, newStart, newCount)
|
|
return header + "\n" + strings.Join(body, "\n") + "\n… local snapshot truncated"
|
|
}
|
|
|
|
func truncateAIInline(value string, maxBytes int) string {
|
|
if len(value) <= maxBytes {
|
|
return value
|
|
}
|
|
runes := []rune(value)
|
|
for len(runes) > 0 && len(string(runes)) > maxBytes-3 {
|
|
runes = runes[:len(runes)-1]
|
|
}
|
|
return string(runes) + "…"
|
|
}
|
|
|
|
// prepareAIDiffFromDetails gives validation a conservative fallback. Thread
|
|
// hunks are the only diff material retained in PRDetails; no complete source is
|
|
// persisted in local AI state.
|
|
func prepareAIDiffFromDetails(pr PRDetails) ([]aiDiffFile, []string, int) {
|
|
var files []aiDiffFile
|
|
for _, thread := range pr.Threads {
|
|
if len(thread.Comments) == 0 || strings.TrimSpace(thread.Comments[0].DiffHunk) == "" {
|
|
continue
|
|
}
|
|
files = append(files, aiDiffFile{
|
|
Path: thread.Path, Text: thread.Comments[0].DiffHunk,
|
|
ChangedLines: map[int]bool{thread.Line: true},
|
|
DeletedLines: map[int]bool{thread.Line: true},
|
|
})
|
|
}
|
|
return files, nil, 0
|
|
}
|
|
|
|
func (s *AIStore) SetResolved(pr PRDetails, threadID string, resolved bool) (ReviewThread, error) {
|
|
state, err := s.Load(pr)
|
|
if err != nil {
|
|
return ReviewThread{}, err
|
|
}
|
|
for i := range state.Threads {
|
|
if state.Threads[i].ID == threadID {
|
|
state.Threads[i].IsResolved = resolved
|
|
if err := s.Save(pr, state); err != nil {
|
|
return ReviewThread{}, err
|
|
}
|
|
result := state.Threads[i]
|
|
result.IsOutdated = result.HeadOID != "" && result.HeadOID != pr.HeadOID
|
|
return result, nil
|
|
}
|
|
}
|
|
return ReviewThread{}, errors.New("local AI thread was not found")
|
|
}
|