Files
diple/ai_store.go
2026-07-29 09:48:50 +02:00

488 lines
14 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...)
}
}
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")
}