diff --git a/MR_config/local/application.example.yml b/MR_config/local/application.example.yml index f575d586..4baef5dd 100644 --- a/MR_config/local/application.example.yml +++ b/MR_config/local/application.example.yml @@ -50,9 +50,11 @@ oauth: # AI 도입 및 AWS S3 공통 구조 external: ai: - base-url: ${AI_BASE_URL:} - api-key: ${AI_API_KEY:} - model: ${AI_MODEL:} + base-url: ${AI_BASE_URL:https://generativelanguage.googleapis.com} + api-key: ${GEMINI_API_KEY:${AI_API_KEY:}} + model: ${AI_MODEL:gemini-3-flash-preview} + connect-timeout: ${AI_CONNECT_TIMEOUT:5s} + read-timeout: ${AI_READ_TIMEOUT:120s} aws: s3: @@ -63,11 +65,11 @@ aws: # 내부 AI 분석 서버 설정 ai: internal: - base-url: ${AI_INTERNAL_BASE_URL} - connect-timeout: ${AI_INTERNAL_CONNECT_TIMEOUT} - read-timeout: ${AI_INTERNAL_READ_TIMEOUT} + base-url: ${AI_INTERNAL_BASE_URL:https://ai.musereview.site} + connect-timeout: ${AI_INTERNAL_CONNECT_TIMEOUT:5s} + read-timeout: ${AI_INTERNAL_READ_TIMEOUT:60s} endpoints: - analyze: ${AI_INTERNAL_ANALYZE_ENDPOINT} + analyze: ${AI_INTERNAL_ANALYZE_ENDPOINT:/analyze} # 프로필 관련 공통 설정 app: diff --git a/src/main/java/com/mr/Application.java b/src/main/java/com/mr/Application.java index d6846c64..87e38f6f 100644 --- a/src/main/java/com/mr/Application.java +++ b/src/main/java/com/mr/Application.java @@ -4,8 +4,10 @@ import org.springframework.boot.autoconfigure.SpringBootApplication; import org.springframework.data.jpa.repository.config.EnableJpaAuditing; import org.springframework.retry.annotation.EnableRetry; +import org.springframework.scheduling.annotation.EnableScheduling; @EnableRetry +@EnableScheduling @EnableJpaAuditing @SpringBootApplication public class Application { diff --git a/src/main/java/com/mr/domain/analysis/controller/AnalysisController.java b/src/main/java/com/mr/domain/analysis/controller/AnalysisController.java index 107da8f3..2bd5ae97 100644 --- a/src/main/java/com/mr/domain/analysis/controller/AnalysisController.java +++ b/src/main/java/com/mr/domain/analysis/controller/AnalysisController.java @@ -1,24 +1,48 @@ package com.mr.domain.analysis.controller; +import com.mr.domain.analysis.dto.req.AnalysisCreateRequestDTO; +import com.mr.domain.analysis.dto.res.AnalysisCreateResponseDTO; import com.mr.domain.analysis.dto.res.AnalysisResultResponseDTO; import com.mr.domain.analysis.dto.res.AnalysisStatusResponseDTO; import com.mr.domain.analysis.service.AnalysisService; import com.mr.global.apipayload.ApiResponse; import com.mr.global.security.SecurityUtil; +import io.swagger.v3.oas.annotations.Operation; +import io.swagger.v3.oas.annotations.tags.Tag; +import jakarta.validation.Valid; import lombok.RequiredArgsConstructor; import org.springframework.web.bind.annotation.GetMapping; import org.springframework.web.bind.annotation.PathVariable; +import org.springframework.web.bind.annotation.PostMapping; +import org.springframework.web.bind.annotation.RequestBody; import org.springframework.web.bind.annotation.RequestMapping; import org.springframework.web.bind.annotation.RestController; @RestController @RequiredArgsConstructor @RequestMapping("/api/analyses") +@Tag(name = "분석", description = "연주 분석 요청 및 결과 조회 API") public class AnalysisController { private final AnalysisService analysisService; + @PostMapping + @Operation( + summary = "분석 요청 생성 API", + description = "연주 데이터를 기반으로 AI 분석을 비동기로 요청합니다." + ) + public ApiResponse createAnalysis( + @Valid @RequestBody AnalysisCreateRequestDTO request + ) { + Long userId = SecurityUtil.getCurrentUserId(); + return ApiResponse.onSuccess(analysisService.createAnalysis(userId, request)); + } + @GetMapping("/{analysisId}/status") + @Operation( + summary = "분석 상태 조회 API", + description = "분석 요청의 현재 처리 상태를 조회합니다." + ) public ApiResponse getAnalysisStatus( @PathVariable Long analysisId ) { @@ -30,6 +54,10 @@ public ApiResponse getAnalysisStatus( } @GetMapping("/{analysisId}") + @Operation( + summary = "분석 결과 조회 API", + description = "완료된 분석 결과와 생성된 리포트를 조회합니다." + ) public ApiResponse getAnalysisResult( @PathVariable Long analysisId ) { @@ -39,4 +67,4 @@ public ApiResponse getAnalysisResult( analysisService.getAnalysisResult(userId, analysisId) ); } -} \ No newline at end of file +} diff --git a/src/main/java/com/mr/domain/analysis/dto/req/AnalysisCreateRequestDTO.java b/src/main/java/com/mr/domain/analysis/dto/req/AnalysisCreateRequestDTO.java new file mode 100644 index 00000000..16752ed8 --- /dev/null +++ b/src/main/java/com/mr/domain/analysis/dto/req/AnalysisCreateRequestDTO.java @@ -0,0 +1,11 @@ +package com.mr.domain.analysis.dto.req; + +import jakarta.validation.constraints.NotNull; +import jakarta.validation.constraints.Positive; + +public record AnalysisCreateRequestDTO( + @NotNull Long playingId, + @NotNull @Positive Integer startBar, + @NotNull @Positive Integer endBar +) { +} diff --git a/src/main/java/com/mr/domain/analysis/dto/res/AnalysisCreateResponseDTO.java b/src/main/java/com/mr/domain/analysis/dto/res/AnalysisCreateResponseDTO.java new file mode 100644 index 00000000..921c3a97 --- /dev/null +++ b/src/main/java/com/mr/domain/analysis/dto/res/AnalysisCreateResponseDTO.java @@ -0,0 +1,19 @@ +package com.mr.domain.analysis.dto.res; + +import com.mr.domain.analysis.entity.Analysis; +import com.mr.domain.analysis.entity.enums.AnalysisStatus; +import java.time.LocalDateTime; + +public record AnalysisCreateResponseDTO( + Long analysisId, + Long playingId, + AnalysisStatus status, + Integer startBar, + Integer endBar, + LocalDateTime createdAt +) { + public static AnalysisCreateResponseDTO from(Analysis analysis) { + return new AnalysisCreateResponseDTO(analysis.getId(), analysis.getPlaying().getId(), analysis.getStatus(), + analysis.getStartBar(), analysis.getEndBar(), analysis.getCreatedAt()); + } +} diff --git a/src/main/java/com/mr/domain/analysis/dto/res/AnalysisResultResponseDTO.java b/src/main/java/com/mr/domain/analysis/dto/res/AnalysisResultResponseDTO.java index 12fa55b6..49e31c44 100644 --- a/src/main/java/com/mr/domain/analysis/dto/res/AnalysisResultResponseDTO.java +++ b/src/main/java/com/mr/domain/analysis/dto/res/AnalysisResultResponseDTO.java @@ -1,6 +1,10 @@ package com.mr.domain.analysis.dto.res; import com.fasterxml.jackson.databind.JsonNode; +import com.fasterxml.jackson.annotation.JsonProperty; +import com.mr.domain.analysis.entity.enums.AnalysisStatus; +import com.mr.domain.backingTrack.entity.BackingTrack; +import com.mr.domain.playing.entity.Playing; import com.mr.domain.analysis.entity.Analysis; import com.mr.domain.analysis.entity.AnalysisReport; import com.mr.domain.analysis.entity.enums.AnalysisGrade; @@ -9,10 +13,17 @@ import com.mr.domain.analysis.entity.enums.ReportGenerationType; import java.math.BigDecimal; import java.time.LocalDateTime; +import java.util.Locale; public record AnalysisResultResponseDTO( Long analysisId, Long playingId, + String title, + String genre, + String key, + Integer bpm, + LocalDateTime playedAt, + AnalysisStatus status, Integer startBar, Integer endBar, Integer totalScore, @@ -20,7 +31,7 @@ public record AnalysisResultResponseDTO( String summary, DomainScores domainScores, Report report, - JsonNode rawResult, + @JsonProperty("result") JsonNode rawResult, LocalDateTime createdAt, LocalDateTime completedAt ) { @@ -30,9 +41,17 @@ public static AnalysisResultResponseDTO from( AnalysisReport analysisReport, JsonNode rawResult ) { + Playing playing = analysis.getPlaying(); + BackingTrack backingTrack = playing.getBackingTrack(); return new AnalysisResultResponseDTO( analysis.getId(), - analysis.getPlaying().getId(), + playing.getId(), + backingTrack.getTitle(), + backingTrack.getGenre(), + formatKey(backingTrack), + playing.getBpm(), + playing.getEndedAt(), + analysis.getStatus(), analysis.getStartBar(), analysis.getEndBar(), analysis.getTotalScore(), @@ -46,11 +65,19 @@ public static AnalysisResultResponseDTO from( ); } + private static String formatKey(BackingTrack backingTrack) { + String scale = backingTrack.getScaleType().name().toLowerCase(Locale.ROOT); + return backingTrack.getKeySignature() + + " " + + Character.toUpperCase(scale.charAt(0)) + + scale.substring(1); + } + public record DomainScores( - BigDecimal scaleScore, - BigDecimal tensionScore, - BigDecimal progressionScore, - BigDecimal voiceLeadingScore + @JsonProperty("scale") BigDecimal scaleScore, + @JsonProperty("tension") BigDecimal tensionScore, + @JsonProperty("progression") BigDecimal progressionScore, + @JsonProperty("voiceLeading") BigDecimal voiceLeadingScore ) { private static DomainScores from(Analysis analysis) { @@ -93,4 +120,4 @@ private static Report fromNullable(AnalysisReport analysisReport) { ); } } -} \ No newline at end of file +} diff --git a/src/main/java/com/mr/domain/analysis/entity/Analysis.java b/src/main/java/com/mr/domain/analysis/entity/Analysis.java index ada3de8f..6c34ed1d 100644 --- a/src/main/java/com/mr/domain/analysis/entity/Analysis.java +++ b/src/main/java/com/mr/domain/analysis/entity/Analysis.java @@ -11,11 +11,14 @@ import java.math.BigDecimal; import java.time.LocalDateTime; +import java.time.temporal.ChronoUnit; import java.util.Objects; import lombok.AccessLevel; import lombok.Builder; import lombok.Getter; import lombok.NoArgsConstructor; +import org.hibernate.annotations.JdbcTypeCode; +import org.hibernate.type.SqlTypes; @Getter @Entity @@ -65,6 +68,7 @@ public class Analysis extends BaseCreatedEntity { @Column(name = "summary", columnDefinition = "text") private String summary; + @JdbcTypeCode(SqlTypes.JSON) @Column(name = "analysis_request_json", nullable = false, columnDefinition = "json") private String analysisRequestJson; @@ -80,12 +84,16 @@ public class Analysis extends BaseCreatedEntity { @Column(name = "voice_leading_score", precision = 5, scale = 2) private BigDecimal voiceLeadingScore; + @JdbcTypeCode(SqlTypes.JSON) @Column(name = "raw_result_json", columnDefinition = "json") private String rawResultJson; @Column(name = "failed_reason", columnDefinition = "text") private String failedReason; + @Column(name = "processing_started_at") + private LocalDateTime processingStartedAt; + @Column(name = "completed_at") private LocalDateTime completedAt; @@ -164,16 +172,39 @@ private static void validateBarRange(Integer startBar, Integer endBar) { } } - public void startProcessing() { + public LocalDateTime startProcessing(LocalDateTime now) { if (this.status != AnalysisStatus.PENDING) { throw new IllegalStateException("PENDING 상태의 분석만 PROCESSING으로 변경할 수 있습니다."); } this.status = AnalysisStatus.PROCESSING; + this.processingStartedAt = nextProcessingStartedAt(now); + return this.processingStartedAt; + } + + public LocalDateTime restartProcessing(LocalDateTime now) { + if (this.status != AnalysisStatus.PROCESSING) { + throw new IllegalStateException("PROCESSING 상태의 분석만 다시 시작할 수 있습니다."); + } + this.processingStartedAt = nextProcessingStartedAt(now); + return this.processingStartedAt; + } + + public boolean isCurrentProcessing(LocalDateTime expectedProcessingStartedAt) { + return this.status == AnalysisStatus.PROCESSING + && Objects.equals(this.processingStartedAt, expectedProcessingStartedAt); + } + + private LocalDateTime nextProcessingStartedAt(LocalDateTime requestedAt) { + LocalDateTime next = Objects.requireNonNull(requestedAt).truncatedTo(ChronoUnit.MILLIS); + if (this.processingStartedAt != null && !next.isAfter(this.processingStartedAt)) { + return this.processingStartedAt.plus(1, ChronoUnit.MILLIS); + } + return next; } public void complete(Integer totalScore, AnalysisGrade grade, String summary, BigDecimal scaleScore, BigDecimal tensionScore, BigDecimal progressionScore, BigDecimal voiceLeadingScore, - String rawResultJson) { + String rawResultJson, LocalDateTime completedAt) { if (this.status != AnalysisStatus.PROCESSING) { throw new IllegalStateException("PROCESSING 상태의 분석만 완료 처리할 수 있습니다."); } @@ -186,15 +217,15 @@ public void complete(Integer totalScore, AnalysisGrade grade, String summary, Bi this.progressionScore = progressionScore; this.voiceLeadingScore = voiceLeadingScore; this.rawResultJson = rawResultJson; - this.completedAt = LocalDateTime.now(); + this.completedAt = Objects.requireNonNull(completedAt); } - public void fail(String failedReason) { + public void fail(String failedReason, LocalDateTime completedAt) { if (this.status == AnalysisStatus.COMPLETED || this.status == AnalysisStatus.FAILED) { throw new IllegalStateException("이미 완료된 분석은 실패 처리할 수 없습니다."); } this.status = AnalysisStatus.FAILED; this.failedReason = failedReason; - this.completedAt = LocalDateTime.now(); + this.completedAt = Objects.requireNonNull(completedAt); } } diff --git a/src/main/java/com/mr/domain/analysis/event/AnalysisRequestedEvent.java b/src/main/java/com/mr/domain/analysis/event/AnalysisRequestedEvent.java new file mode 100644 index 00000000..50516732 --- /dev/null +++ b/src/main/java/com/mr/domain/analysis/event/AnalysisRequestedEvent.java @@ -0,0 +1,4 @@ +package com.mr.domain.analysis.event; + +public record AnalysisRequestedEvent(Long analysisId) { +} diff --git a/src/main/java/com/mr/domain/analysis/event/listener/AnalysisRequestedEventListener.java b/src/main/java/com/mr/domain/analysis/event/listener/AnalysisRequestedEventListener.java new file mode 100644 index 00000000..077db021 --- /dev/null +++ b/src/main/java/com/mr/domain/analysis/event/listener/AnalysisRequestedEventListener.java @@ -0,0 +1,36 @@ +package com.mr.domain.analysis.event.listener; + +import com.mr.domain.analysis.event.AnalysisRequestedEvent; +import com.mr.domain.analysis.service.AnalysisProcessingService; +import lombok.extern.slf4j.Slf4j; +import org.springframework.beans.factory.annotation.Qualifier; +import org.springframework.core.task.TaskExecutor; +import org.springframework.stereotype.Component; +import org.springframework.transaction.event.TransactionPhase; +import org.springframework.transaction.event.TransactionalEventListener; + +@Component +@Slf4j +public class AnalysisRequestedEventListener { + + private final AnalysisProcessingService analysisProcessingService; + private final TaskExecutor taskExecutor; + + public AnalysisRequestedEventListener( + AnalysisProcessingService analysisProcessingService, + @Qualifier("applicationTaskExecutor") TaskExecutor taskExecutor + ) { + this.analysisProcessingService = analysisProcessingService; + this.taskExecutor = taskExecutor; + } + + @TransactionalEventListener(phase = TransactionPhase.AFTER_COMMIT) + public void handle(AnalysisRequestedEvent event) { + try { + taskExecutor.execute(() -> analysisProcessingService.process(event.analysisId())); + } catch (RuntimeException exception) { + log.warn("AI analysis submission failed; it will be retried by recovery. analysisId={}", + event.analysisId(), exception); + } + } +} diff --git a/src/main/java/com/mr/domain/analysis/exception/AnalysisErrorStatus.java b/src/main/java/com/mr/domain/analysis/exception/AnalysisErrorStatus.java index ed115a24..f9c44b9b 100644 --- a/src/main/java/com/mr/domain/analysis/exception/AnalysisErrorStatus.java +++ b/src/main/java/com/mr/domain/analysis/exception/AnalysisErrorStatus.java @@ -9,22 +9,35 @@ @AllArgsConstructor public enum AnalysisErrorStatus implements BaseCode { - // 400_01, 400_02는 스펙상 마디 검증용으로 예약됨 + INVALID_BAR_ORDER(HttpStatus.BAD_REQUEST, "ANALYSIS_400_01", "분석 시작 마디는 종료 마디보다 클 수 없습니다."), + + INVALID_BAR_RANGE(HttpStatus.BAD_REQUEST, "ANALYSIS_400_02", "분석 가능한 마디 범위를 벗어났습니다."), + ANALYSIS_INVALID_REQUEST(HttpStatus.BAD_REQUEST, "ANALYSIS_400_03", "필수 정보가 누락되었습니다."), ANALYSIS_OWNER_MISMATCH(HttpStatus.BAD_REQUEST, "ANALYSIS_400_04", "playing 소유자와 user가 일치하지 않습니다."), + EMPTY_NOTE_RANGE(HttpStatus.BAD_REQUEST, "ANALYSIS_400_05", "선택한 마디 범위에 분석할 연주 노트가 없습니다."), + ANALYSIS_ACCESS_DENIED(HttpStatus.FORBIDDEN, "ANALYSIS_403_01", "해당 분석 결과에 접근할 수 없습니다."), ANALYSIS_NOT_FOUND(HttpStatus.NOT_FOUND, "ANALYSIS_404_01", "분석 결과를 찾을 수 없습니다."), - ANALYSIS_NOT_COMPLETED(HttpStatus.CONFLICT, "ANALYSIS_409_01", "아직 완료되지 않은 분석입니다."), + ANALYSIS_ALREADY_IN_PROGRESS( + HttpStatus.CONFLICT, + "ANALYSIS_409_01", + "이미 처리 중이거나 완료된 분석 요청이 있습니다." + ), + + ANALYSIS_NOT_COMPLETED(HttpStatus.CONFLICT, "ANALYSIS_409_02", "아직 완료되지 않은 분석입니다."), INVALID_RAW_RESULT(HttpStatus.INTERNAL_SERVER_ERROR, "ANALYSIS_500_01", "저장된 분석 결과를 처리할 수 없습니다."), + INVALID_ANALYSIS_REQUEST(HttpStatus.INTERNAL_SERVER_ERROR, "ANALYSIS_500_02", "저장된 분석 요청을 처리할 수 없습니다."), + ; private final HttpStatus status; private final String code; private final String message; -} \ No newline at end of file +} diff --git a/src/main/java/com/mr/domain/analysis/factory/AnalysisRequestFactory.java b/src/main/java/com/mr/domain/analysis/factory/AnalysisRequestFactory.java new file mode 100644 index 00000000..1cbc227b --- /dev/null +++ b/src/main/java/com/mr/domain/analysis/factory/AnalysisRequestFactory.java @@ -0,0 +1,116 @@ +package com.mr.domain.analysis.factory; + +import com.mr.domain.analysis.exception.AnalysisErrorStatus; +import com.mr.domain.backingTrack.entity.BackingTrack; +import com.mr.domain.backingTrack.entity.ChordProgression; +import com.mr.domain.playing.entity.MidiEventData; +import com.mr.domain.playing.entity.Playing; +import com.mr.domain.playing.entity.enums.MidiType; +import com.mr.global.apipayload.exception.GeneralException; +import com.mr.global.client.ai.AiAnalysisRequest; +import java.util.Arrays; +import java.util.Comparator; +import java.util.List; +import java.util.Locale; +import java.util.concurrent.atomic.AtomicInteger; +import org.springframework.stereotype.Component; + +@Component +public class AnalysisRequestFactory { + + private static final int MAX_BAR_COUNT = 32; + private static final double MILLIS_PER_MINUTE = 60_000D; + + public AiAnalysisRequest create(Playing playing, int startBar, int endBar) { + BackingTrack track = playing.getBackingTrack(); + int[] timeSignature = parseTimeSignature(track.getTimeSignature()); + double barDurationMs = MILLIS_PER_MINUTE / playing.getBpm() + * timeSignature[0] * 4D / timeSignature[1]; + + validateBarRange(track, startBar, endBar, barDurationMs); + + double startOffsetMs = (startBar - 1) * barDurationMs; + double endOffsetMs = endBar * barDurationMs; + AtomicInteger noteIndex = new AtomicInteger(); + + List notes = playing.getMidiData().stream() + .filter(event -> event.getTimestampMs() >= startOffsetMs + && event.getTimestampMs() < endOffsetMs) + .sorted(Comparator.comparingLong(MidiEventData::getTimestampMs) + .thenComparingInt(MidiEventData::getSequence)) + .map(event -> new AiAnalysisRequest.Note( + noteIndex.getAndIncrement(), + toNoteType(event.getType()), + event.getPitch(), + event.getVelocity(), + event.getTimestampMs() - startOffsetMs + )) + .toList(); + if (notes.isEmpty()) { + throw new GeneralException(AnalysisErrorStatus.EMPTY_NOTE_RANGE); + } + + List chords = track.getChordProgressions().stream() + .filter(chord -> chord.getMeasureNo() >= startBar && chord.getMeasureNo() <= endBar) + .sorted(Comparator.comparingInt(ChordProgression::getMeasureNo) + .thenComparingInt(ChordProgression::getSequenceNo)) + .map(chord -> new AiAnalysisRequest.Chord( + chord.getMeasureNo() - startBar + 1, + chord.getSequenceNo().doubleValue(), + chord.getChordName() + )) + .toList(); + + AiAnalysisRequest.Meta meta = new AiAnalysisRequest.Meta( + playing.getBpm().doubleValue(), + Arrays.stream(timeSignature).boxed().toList(), + new AiAnalysisRequest.Key( + track.getKeySignature(), + track.getScaleType().name().toLowerCase(Locale.ROOT) + ), + track.getGenre().toLowerCase(Locale.ROOT), + track.getLevel().name().toLowerCase(Locale.ROOT) + ); + return new AiAnalysisRequest(meta, chords, notes); + } + + private AiAnalysisRequest.NoteType toNoteType(MidiType type) { + return switch (type) { + case NOTE_ON -> AiAnalysisRequest.NoteType.NOTE_ON; + case NOTE_OFF -> AiAnalysisRequest.NoteType.NOTE_OFF; + }; + } + + private void validateBarRange( + BackingTrack track, + int startBar, + int endBar, + double barDurationMs + ) { + if (startBar > endBar) { + throw new GeneralException(AnalysisErrorStatus.INVALID_BAR_ORDER); + } + int barCount = endBar - startBar + 1; + int totalBars = (int) Math.ceil(track.getPlaytimeSec() * 1_000D / barDurationMs); + if (barCount > MAX_BAR_COUNT || endBar > totalBars) { + throw new GeneralException(AnalysisErrorStatus.INVALID_BAR_RANGE); + } + } + + private int[] parseTimeSignature(String value) { + try { + String[] parts = value.split("/"); + if (parts.length != 2) { + throw new NumberFormatException(); + } + int numerator = Integer.parseInt(parts[0]); + int denominator = Integer.parseInt(parts[1]); + if (numerator <= 0 || denominator <= 0) { + throw new NumberFormatException(); + } + return new int[]{numerator, denominator}; + } catch (NumberFormatException exception) { + throw new GeneralException(AnalysisErrorStatus.ANALYSIS_INVALID_REQUEST); + } + } +} diff --git a/src/main/java/com/mr/domain/analysis/generator/AnalysisResultEnricher.java b/src/main/java/com/mr/domain/analysis/generator/AnalysisResultEnricher.java new file mode 100644 index 00000000..0cc41231 --- /dev/null +++ b/src/main/java/com/mr/domain/analysis/generator/AnalysisResultEnricher.java @@ -0,0 +1,46 @@ +package com.mr.domain.analysis.generator; + +import com.fasterxml.jackson.databind.JsonNode; +import com.fasterxml.jackson.databind.node.ObjectNode; +import java.math.BigDecimal; +import java.util.ArrayList; +import java.util.Comparator; +import org.springframework.stereotype.Component; + +@Component +public class AnalysisResultEnricher { + + public JsonNode enrich(JsonNode result) { + String existingSummary = result.path("summary").asText(); + if (!existingSummary.isBlank()) { + return result; + } + + ObjectNode enriched = ((ObjectNode) result).deepCopy(); + enriched.put("summary", generateSummary(result)); + return enriched; + } + + private String generateSummary(JsonNode result) { + JsonNode scores = result.path("scores"); + var domainScores = new ArrayList(); + scores.path("domains").fields().forEachRemaining(entry -> { + if (entry.getValue().isNumber()) { + domainScores.add(new DomainScore(entry.getKey(), entry.getValue().decimalValue())); + } + }); + DomainScore strongest = domainScores.stream() + .max(Comparator.comparing(DomainScore::score)) + .orElseThrow(() -> new IllegalStateException("도메인 점수가 없습니다.")); + return switch (strongest.name()) { + case "스케일" -> "조성에 어울리는 음 선택이 안정적으로 이어졌어요."; + case "텐션" -> "긴장감을 살리는 텐션 활용이 자연스럽게 이어졌어요."; + case "진행" -> "코드 진행의 흐름을 안정적으로 잘 따라갔어요."; + case "코드 연결" -> "코드 사이의 음 연결이 자연스럽고 매끄러웠어요."; + default -> strongest.name() + " 영역의 흐름이 전반적으로 안정적이었어요."; + }; + } + + private record DomainScore(String name, BigDecimal score) { + } +} diff --git a/src/main/java/com/mr/domain/analysis/generator/RuleBasedReportGenerator.java b/src/main/java/com/mr/domain/analysis/generator/RuleBasedReportGenerator.java new file mode 100644 index 00000000..d26f264f --- /dev/null +++ b/src/main/java/com/mr/domain/analysis/generator/RuleBasedReportGenerator.java @@ -0,0 +1,165 @@ +package com.mr.domain.analysis.generator; + +import com.fasterxml.jackson.databind.JsonNode; +import java.math.BigDecimal; +import java.math.RoundingMode; +import java.util.LinkedHashSet; +import java.util.Set; +import org.springframework.stereotype.Component; + +@Component +public class RuleBasedReportGenerator { + + public String generate(JsonNode result) { + JsonNode meta = result.path("meta"); + JsonNode scores = result.path("scores"); + JsonNode domains = scores.path("domains"); + + StringBuilder report = new StringBuilder(); + report.append("# 연주 분석 리포트\n\n"); + report.append("**조성** ").append(text(meta, "key", "-")) + .append(" · **장르** ").append(text(meta, "genre", "-")) + .append(" · **박자** ").append(timeSignature(meta.path("time_signature"))) + .append(" · **템포** ").append(number(meta.path("bpm"))).append(" bpm\n\n"); + + report.append("## 총평\n\n") + .append("전체 점수는 **").append(number(scores.path("final_score"))) + .append(" / 100**입니다. ") + .append(strongestAndWeakest(domains)).append(" ") + .append(text(result, "summary", "영역별 점수를 기준으로 강점과 보완점을 정리했습니다.")) + .append("\n\n"); + + report.append("## 잘한 점\n\n") + .append("- ").append(strongestDomain(domains)) + .append("으로 네 영역 중 가장 높은 평가를 받았습니다.\n"); + appendCoverage(report, scores.path("coverage")); + report.append("\n"); + + report.append("## 진행 맥락\n\n"); + JsonNode rules = result.path("harmonic_rules"); + if (rules.isArray() && !rules.isEmpty()) { + Set labels = new LinkedHashSet<>(); + for (JsonNode rule : rules) { + labels.add(text(rule, "label", text(rule, "rule", "감지된 화성 진행"))); + } + labels.stream().limit(5).forEach(label -> report.append("- ").append(label).append("\n")); + } else { + report.append("- 감지된 화성 진행을 바탕으로 연주를 분석했습니다.\n"); + } + + report.append("\n## 개선 제안\n\n") + .append("- ").append(weakestDomain(domains)) + .append(" 영역이 상대적으로 낮습니다. 백킹트랙 속도를 낮추고 같은 구간을 반복해 보세요.\n"); + appendScaleSuggestion(report, scores.path("scale_appropriateness")); + appendTimingSuggestion(report, result.path("timing_deviations")); + appendLearningRecommendations(report, result.path("learning_recommendations")); + + report.append("\n## 점수 요약\n\n") + .append("- 종합 점수: ").append(number(scores.path("final_score"))).append(" / 100\n"); + domains.fields().forEachRemaining(entry -> + report.append("- ").append(entry.getKey()).append(": ") + .append(number(entry.getValue())).append("\n")); + return report.toString(); + } + + private String strongestAndWeakest(JsonNode domains) { + return strongestDomain(domains) + "이 강점이며, " + weakestDomain(domains) + "을 우선 보완하면 좋습니다."; + } + + private String strongestDomain(JsonNode domains) { + return domainAtExtreme(domains, true); + } + + private String weakestDomain(JsonNode domains) { + return domainAtExtreme(domains, false); + } + + private String domainAtExtreme(JsonNode domains, boolean maximum) { + String selected = "연주"; + BigDecimal selectedScore = null; + var fields = domains.fields(); + while (fields.hasNext()) { + var field = fields.next(); + if (!field.getValue().isNumber()) { + continue; + } + BigDecimal score = field.getValue().decimalValue(); + if (selectedScore == null + || (maximum && score.compareTo(selectedScore) > 0) + || (!maximum && score.compareTo(selectedScore) < 0)) { + selected = field.getKey(); + selectedScore = score; + } + } + return selected + (selectedScore == null ? "" : " (" + format(selectedScore) + "점)"); + } + + private void appendTimingSuggestion(StringBuilder report, JsonNode timing) { + int flaggedCount = timing.path("summary").path("flagged_count").asInt(0); + if (flaggedCount > 0) { + report.append("- 박자가 흔들린 음이 ").append(flaggedCount) + .append("개 감지되었습니다. 메트로놈과 함께 해당 구간을 점검해 보세요.\n"); + } + } + + private void appendCoverage(StringBuilder report, JsonNode coverage) { + String note = coverage.path("note").asText(); + if (!note.isBlank()) { + report.append("- 분석 범위: ").append(note).append("\n"); + } + } + + private void appendScaleSuggestion(StringBuilder report, JsonNode scaleAppropriateness) { + JsonNode notes = scaleAppropriateness.path("out_of_scale_notes"); + if (!notes.isArray() || notes.isEmpty()) { + return; + } + JsonNode first = notes.path(0); + report.append("- 스케일 밖 음이 ").append(notes.size()).append("개 감지되었습니다."); + String suggestion = first.path("suggestion").asText(); + if (!suggestion.isBlank()) { + report.append(" 예: ").append(suggestion); + } + report.append("\n"); + } + + private void appendLearningRecommendations(StringBuilder report, JsonNode recommendations) { + if (!recommendations.isArray()) { + return; + } + int count = Math.min(recommendations.size(), 2); + for (int index = 0; index < count; index++) { + JsonNode recommendation = recommendations.path(index); + String title = text(recommendation, "title", "추천 학습"); + String reason = recommendation.path("reason").asText(); + String studyTip = recommendation.path("study_tip").asText(); + report.append("- ").append(title).append(": "); + if (!reason.isBlank()) { + report.append(reason).append(" "); + } + if (!studyTip.isBlank()) { + report.append(studyTip); + } + report.append("\n"); + } + } + + private String timeSignature(JsonNode node) { + return node.isArray() && node.size() == 2 + ? node.path(0).asText() + "/" + node.path(1).asText() + : "-"; + } + + private String text(JsonNode parent, String field, String fallback) { + String value = parent.path(field).asText(); + return value.isBlank() ? fallback : value; + } + + private String number(JsonNode node) { + return node.isNumber() ? format(node.decimalValue()) : "-"; + } + + private String format(BigDecimal value) { + return value.setScale(1, RoundingMode.HALF_UP).stripTrailingZeros().toPlainString(); + } +} diff --git a/src/main/java/com/mr/domain/analysis/model/AnalysisProcessingClaim.java b/src/main/java/com/mr/domain/analysis/model/AnalysisProcessingClaim.java new file mode 100644 index 00000000..c66fff6f --- /dev/null +++ b/src/main/java/com/mr/domain/analysis/model/AnalysisProcessingClaim.java @@ -0,0 +1,9 @@ +package com.mr.domain.analysis.model; + +import java.time.LocalDateTime; + +public record AnalysisProcessingClaim( + String requestJson, + LocalDateTime processingStartedAt +) { +} diff --git a/src/main/java/com/mr/domain/analysis/model/GeneratedAnalysisReport.java b/src/main/java/com/mr/domain/analysis/model/GeneratedAnalysisReport.java new file mode 100644 index 00000000..1ea3872b --- /dev/null +++ b/src/main/java/com/mr/domain/analysis/model/GeneratedAnalysisReport.java @@ -0,0 +1,12 @@ +package com.mr.domain.analysis.model; + +import com.mr.domain.analysis.entity.enums.ReportGenerationType; + +public record GeneratedAnalysisReport( + ReportGenerationType generationType, + String content, + String modelName, + String promptVersion, + LlmCallMetadata llmCall +) { +} diff --git a/src/main/java/com/mr/domain/analysis/model/LlmCallMetadata.java b/src/main/java/com/mr/domain/analysis/model/LlmCallMetadata.java new file mode 100644 index 00000000..48bd7b1a --- /dev/null +++ b/src/main/java/com/mr/domain/analysis/model/LlmCallMetadata.java @@ -0,0 +1,21 @@ +package com.mr.domain.analysis.model; + +import com.fasterxml.jackson.databind.JsonNode; +import com.mr.domain.mentor.entity.enums.LlmCallStatus; +import java.math.BigDecimal; + +public record LlmCallMetadata( + LlmCallStatus status, + String modelName, + String promptVersion, + JsonNode promptSnapshot, + Integer promptTokens, + Integer completionTokens, + Integer totalTokens, + BigDecimal temperature, + Integer latencyMs, + boolean cacheHit, + String inputHash, + String errorMessage +) { +} diff --git a/src/main/java/com/mr/domain/analysis/repository/AnalysisRepository.java b/src/main/java/com/mr/domain/analysis/repository/AnalysisRepository.java index 6c451df3..6d7915e3 100644 --- a/src/main/java/com/mr/domain/analysis/repository/AnalysisRepository.java +++ b/src/main/java/com/mr/domain/analysis/repository/AnalysisRepository.java @@ -2,14 +2,46 @@ import com.mr.domain.analysis.entity.Analysis; import com.mr.domain.analysis.entity.enums.AnalysisStatus; +import jakarta.persistence.LockModeType; import java.time.LocalDateTime; import java.util.List; +import java.util.Optional; +import org.springframework.data.domain.Pageable; import org.springframework.data.jpa.repository.JpaRepository; +import org.springframework.data.jpa.repository.Lock; import org.springframework.data.jpa.repository.Query; import org.springframework.data.repository.query.Param; public interface AnalysisRepository extends JpaRepository { + boolean existsByPlayingIdAndStatusIn(Long playingId, List statuses); + + @Lock(LockModeType.PESSIMISTIC_WRITE) + @Query("select a from Analysis a where a.id = :analysisId") + Optional findByIdForUpdate(@Param("analysisId") Long analysisId); + + @Query(""" + select a.id from Analysis a + where a.status = :status and a.createdAt <= :cutoff + order by a.createdAt asc, a.id asc + """) + List findIdsByStatusAndCreatedAtBefore( + @Param("status") AnalysisStatus status, + @Param("cutoff") LocalDateTime cutoff, + Pageable pageable + ); + + @Query(""" + select a.id from Analysis a + where a.status = :status and a.processingStartedAt <= :cutoff + order by a.processingStartedAt asc, a.id asc + """) + List findIdsByStatusAndProcessingStartedAtBefore( + @Param("status") AnalysisStatus status, + @Param("cutoff") LocalDateTime cutoff, + Pageable pageable + ); + @Query(""" select a from Analysis a where a.playing.id in :playingIds and a.status = :status @@ -26,7 +58,6 @@ List findByPlayingIdInAndStatusOrderByCreatedAtDescIdDesc( List findByPlayingIdAndUserIdOrderByStartBarAscIdAsc( @Param("playingId") Long playingId, @Param("userId") Long userId); - // 통계 집계용 - completedAt 기준(Analysis.complete()에서 세팅되는 값) @Query(""" select a from Analysis a where a.user.userId = :userId @@ -40,7 +71,6 @@ List findByUserAndStatusSince( @Param("since") LocalDateTime since ); - // 통계 집계용 전체 기간 요약 - row를 끌어오지 않고 단일 행으로 집계 @Query(""" select count(a) as analysisCount, avg(a.totalScore) as averageTotalScore from Analysis a diff --git a/src/main/java/com/mr/domain/analysis/scheduler/AnalysisRecoveryScheduler.java b/src/main/java/com/mr/domain/analysis/scheduler/AnalysisRecoveryScheduler.java new file mode 100644 index 00000000..7bbd0e27 --- /dev/null +++ b/src/main/java/com/mr/domain/analysis/scheduler/AnalysisRecoveryScheduler.java @@ -0,0 +1,92 @@ +package com.mr.domain.analysis.scheduler; + +import com.mr.domain.analysis.entity.enums.AnalysisStatus; +import com.mr.domain.analysis.repository.AnalysisRepository; +import java.time.Clock; +import java.time.Duration; +import java.time.LocalDateTime; +import java.util.List; + +import com.mr.domain.analysis.service.AnalysisProcessingService; +import lombok.extern.slf4j.Slf4j; +import org.springframework.beans.factory.annotation.Qualifier; +import org.springframework.beans.factory.annotation.Value; +import org.springframework.core.task.TaskExecutor; +import org.springframework.data.domain.PageRequest; +import org.springframework.scheduling.annotation.Scheduled; +import org.springframework.stereotype.Component; + +@Slf4j +@Component +public class AnalysisRecoveryScheduler { + + private final AnalysisRepository analysisRepository; + private final AnalysisProcessingService analysisProcessingService; + private final TaskExecutor taskExecutor; + private final Clock clock; + private final Duration pendingThreshold; + private final Duration processingThreshold; + private final int batchSize; + + public AnalysisRecoveryScheduler( + AnalysisRepository analysisRepository, + AnalysisProcessingService analysisProcessingService, + @Qualifier("applicationTaskExecutor") TaskExecutor taskExecutor, + Clock clock, + @Value("${analysis.recovery.pending-threshold:1m}") Duration pendingThreshold, + @Value("${analysis.recovery.processing-threshold:5m}") Duration processingThreshold, + @Value("${analysis.recovery.batch-size:20}") int batchSize + ) { + this.analysisRepository = analysisRepository; + this.analysisProcessingService = analysisProcessingService; + this.taskExecutor = taskExecutor; + this.clock = clock; + this.pendingThreshold = pendingThreshold; + this.processingThreshold = processingThreshold; + this.batchSize = batchSize; + } + + @Scheduled( + initialDelayString = "${analysis.recovery.initial-delay-ms:10000}", + fixedDelayString = "${analysis.recovery.fixed-delay-ms:30000}" + ) + public void recoverPendingAnalyses() { + LocalDateTime now = LocalDateTime.now(clock); + List pendingIds = analysisRepository.findIdsByStatusAndCreatedAtBefore( + AnalysisStatus.PENDING, + now.minus(pendingThreshold), + PageRequest.of(0, batchSize) + ); + submitPending(pendingIds); + + LocalDateTime processingCutoff = now.minus(processingThreshold); + List processingIds = analysisRepository.findIdsByStatusAndProcessingStartedAtBefore( + AnalysisStatus.PROCESSING, + processingCutoff, + PageRequest.of(0, batchSize) + ); + submitStaleProcessing(processingIds, processingCutoff); + } + + private void submitPending(List analysisIds) { + for (Long analysisId : analysisIds) { + try { + taskExecutor.execute(() -> analysisProcessingService.process(analysisId)); + } catch (RuntimeException exception) { + log.warn("Recovered AI analysis submission failed; it will be retried. analysisId={}", + analysisId, exception); + } + } + } + + private void submitStaleProcessing(List analysisIds, LocalDateTime cutoff) { + for (Long analysisId : analysisIds) { + try { + taskExecutor.execute(() -> analysisProcessingService.recoverStaleProcessing(analysisId, cutoff)); + } catch (RuntimeException exception) { + log.warn("Stale AI analysis submission failed; it will be retried. analysisId={}", + analysisId, exception); + } + } + } +} diff --git a/src/main/java/com/mr/domain/analysis/service/AnalysisProcessingService.java b/src/main/java/com/mr/domain/analysis/service/AnalysisProcessingService.java new file mode 100644 index 00000000..4733f46f --- /dev/null +++ b/src/main/java/com/mr/domain/analysis/service/AnalysisProcessingService.java @@ -0,0 +1,77 @@ +package com.mr.domain.analysis.service; + +import com.fasterxml.jackson.databind.JsonNode; +import com.fasterxml.jackson.databind.ObjectMapper; +import com.mr.domain.analysis.generator.AnalysisResultEnricher; +import com.mr.domain.analysis.model.AnalysisProcessingClaim; +import com.mr.domain.analysis.model.GeneratedAnalysisReport; +import com.mr.global.client.ai.AiAnalysisRequest; +import com.mr.global.client.ai.AiServerClient; +import java.time.LocalDateTime; +import java.util.Optional; +import java.util.function.Supplier; +import lombok.RequiredArgsConstructor; +import lombok.extern.slf4j.Slf4j; +import org.springframework.stereotype.Service; + +@Slf4j +@Service +@RequiredArgsConstructor +public class AnalysisProcessingService { + + private final AnalysisStateService analysisStateService; + private final AiServerClient aiServerClient; + private final ReportGenerationService reportGenerationService; + private final AnalysisResultEnricher analysisResultEnricher; + private final ObjectMapper objectMapper; + + public void process(Long analysisId) { + processClaimed(analysisId, () -> analysisStateService.startProcessing(analysisId)); + } + + public void recoverStaleProcessing(Long analysisId, LocalDateTime cutoff) { + processClaimed(analysisId, () -> analysisStateService.restartStaleProcessing(analysisId, cutoff)); + } + + private void processClaimed(Long analysisId, Supplier> claim) { + AnalysisProcessingClaim activeClaim = null; + try { + Optional claimed = claim.get(); + if (claimed.isEmpty()) { + return; + } + activeClaim = claimed.get(); + AiAnalysisRequest request = objectMapper.readValue(activeClaim.requestJson(), AiAnalysisRequest.class); + JsonNode result = aiServerClient.requestAnalysis(request); + analysisStateService.validateResult(result); + result = analysisResultEnricher.enrich(result); + GeneratedAnalysisReport report = reportGenerationService.generate(result); + boolean completed = analysisStateService.complete( + analysisId, + activeClaim.processingStartedAt(), + result, + objectMapper.writeValueAsString(result), + report + ); + if (!completed) { + log.info("Ignoring stale AI analysis completion. analysisId={}", analysisId); + } + } catch (Exception exception) { + log.error("AI analysis failed. analysisId={}", analysisId, exception); + if (activeClaim != null) { + try { + boolean failed = analysisStateService.fail( + analysisId, + activeClaim.processingStartedAt(), + exception.getMessage() + ); + if (!failed) { + log.info("Ignoring stale AI analysis failure. analysisId={}", analysisId); + } + } catch (Exception failException) { + log.error("Failed to persist AI analysis failure. analysisId={}", analysisId, failException); + } + } + } + } +} diff --git a/src/main/java/com/mr/domain/analysis/service/AnalysisService.java b/src/main/java/com/mr/domain/analysis/service/AnalysisService.java index 6ee09412..13d0b7ad 100644 --- a/src/main/java/com/mr/domain/analysis/service/AnalysisService.java +++ b/src/main/java/com/mr/domain/analysis/service/AnalysisService.java @@ -3,19 +3,29 @@ import com.fasterxml.jackson.core.JsonProcessingException; import com.fasterxml.jackson.databind.JsonNode; import com.fasterxml.jackson.databind.ObjectMapper; +import com.mr.domain.analysis.dto.req.AnalysisCreateRequestDTO; +import com.mr.domain.analysis.dto.res.AnalysisCreateResponseDTO; import com.mr.domain.analysis.dto.res.AnalysisResultResponseDTO; import com.mr.domain.analysis.dto.res.AnalysisStatusResponseDTO; import com.mr.domain.analysis.entity.Analysis; import com.mr.domain.analysis.entity.AnalysisReport; import com.mr.domain.analysis.entity.enums.AnalysisStatus; import com.mr.domain.analysis.entity.enums.LlmStatus; +import com.mr.domain.analysis.event.AnalysisRequestedEvent; import com.mr.domain.analysis.exception.AnalysisErrorStatus; +import com.mr.domain.analysis.factory.AnalysisRequestFactory; import com.mr.domain.analysis.repository.AnalysisReportRepository; import com.mr.domain.analysis.repository.AnalysisRepository; +import com.mr.domain.playing.entity.Playing; +import com.mr.domain.playing.entity.enums.PlayingStatus; +import com.mr.domain.playing.exception.PlayingErrorStatus; +import com.mr.domain.playing.repository.PlayingRepository; import com.mr.global.apipayload.exception.GeneralException; +import java.util.List; import java.util.Objects; import lombok.RequiredArgsConstructor; import lombok.extern.slf4j.Slf4j; +import org.springframework.context.ApplicationEventPublisher; import org.springframework.stereotype.Service; import org.springframework.transaction.annotation.Transactional; @@ -27,8 +37,51 @@ public class AnalysisService { private final AnalysisRepository analysisRepository; private final AnalysisReportRepository analysisReportRepository; + private final PlayingRepository playingRepository; + private final AnalysisRequestFactory analysisRequestFactory; + private final ApplicationEventPublisher eventPublisher; private final ObjectMapper objectMapper; + @Transactional + public AnalysisCreateResponseDTO createAnalysis( + Long userId, + AnalysisCreateRequestDTO request + ) { + Playing playing = playingRepository.findByIdWithBackingTrackForUpdate(request.playingId()) + .orElseThrow(() -> new GeneralException(PlayingErrorStatus.PLAYING_NOT_FOUND)); + if (!Objects.equals(playing.getUser().getUserId(), userId)) { + throw new GeneralException(PlayingErrorStatus.PLAYING_ACCESS_DENIED); + } + if (playing.getStatus() != PlayingStatus.COMPLETED) { + throw new GeneralException(PlayingErrorStatus.INVALID_PLAYING_STATUS); + } + if (analysisRepository.existsByPlayingIdAndStatusIn( + playing.getId(), + List.of( + AnalysisStatus.PENDING, + AnalysisStatus.PROCESSING, + AnalysisStatus.COMPLETED + ) + )) { + throw new GeneralException(AnalysisErrorStatus.ANALYSIS_ALREADY_IN_PROGRESS); + } + + String requestJson; + try { + requestJson = objectMapper.writeValueAsString( + analysisRequestFactory.create(playing, request.startBar(), request.endBar()) + ); + } catch (JsonProcessingException exception) { + throw new GeneralException(AnalysisErrorStatus.ANALYSIS_INVALID_REQUEST); + } + + Analysis analysis = analysisRepository.save(Analysis.createPending( + playing.getUser(), playing, request.startBar(), request.endBar(), requestJson + )); + eventPublisher.publishEvent(new AnalysisRequestedEvent(analysis.getId())); + return AnalysisCreateResponseDTO.from(analysis); + } + public AnalysisStatusResponseDTO getAnalysisStatus( Long userId, Long analysisId @@ -128,4 +181,4 @@ private JsonNode parseRawResult( ); } } -} \ No newline at end of file +} diff --git a/src/main/java/com/mr/domain/analysis/service/AnalysisStateService.java b/src/main/java/com/mr/domain/analysis/service/AnalysisStateService.java new file mode 100644 index 00000000..af7d13c1 --- /dev/null +++ b/src/main/java/com/mr/domain/analysis/service/AnalysisStateService.java @@ -0,0 +1,229 @@ +package com.mr.domain.analysis.service; + +import com.fasterxml.jackson.databind.JsonNode; +import com.mr.domain.analysis.entity.Analysis; +import com.mr.domain.analysis.entity.AnalysisReport; +import com.mr.domain.analysis.entity.enums.AnalysisGrade; +import com.mr.domain.analysis.entity.enums.AnalysisStatus; +import com.mr.domain.analysis.entity.enums.ReportGenerationType; +import com.mr.domain.analysis.exception.AnalysisErrorStatus; +import com.mr.domain.analysis.model.AnalysisProcessingClaim; +import com.mr.domain.analysis.model.GeneratedAnalysisReport; +import com.mr.domain.analysis.model.LlmCallMetadata; +import com.mr.domain.analysis.repository.AnalysisReportRepository; +import com.mr.domain.analysis.repository.AnalysisRepository; +import com.mr.domain.mentor.entity.LlmCallLog; +import com.mr.domain.mentor.entity.enums.LlmCallStatus; +import com.mr.domain.mentor.entity.enums.LlmPurpose; +import com.mr.domain.mentor.repository.LlmCallLogRepository; +import com.mr.global.apipayload.exception.GeneralException; +import com.mr.global.event.AnalysisCompletedEvent; +import java.math.BigDecimal; +import java.math.RoundingMode; +import java.time.Clock; +import java.time.LocalDateTime; +import java.util.Optional; +import lombok.RequiredArgsConstructor; +import org.springframework.context.ApplicationEventPublisher; +import org.springframework.stereotype.Service; +import org.springframework.transaction.annotation.Transactional; + +@Service +@RequiredArgsConstructor +public class AnalysisStateService { + + private static final BigDecimal MIN_SCORE = BigDecimal.ZERO; + private static final BigDecimal MAX_SCORE = new BigDecimal("100"); + private static final String SCALE = "스케일"; + private static final String TENSION = "텐션"; + private static final String PROGRESSION = "진행"; + private static final String VOICE_LEADING = "코드 연결"; + + private final AnalysisRepository analysisRepository; + private final AnalysisReportRepository analysisReportRepository; + private final LlmCallLogRepository llmCallLogRepository; + private final ApplicationEventPublisher eventPublisher; + private final Clock clock; + + @Transactional + public Optional startProcessing(Long analysisId) { + Analysis analysis = analysisRepository.findByIdForUpdate(analysisId) + .orElseThrow(() -> new GeneralException(AnalysisErrorStatus.ANALYSIS_NOT_FOUND)); + if (analysis.getStatus() != AnalysisStatus.PENDING) { + return Optional.empty(); + } + Optional requestJson = requestJsonOrFail(analysis); + if (requestJson.isEmpty()) { + return Optional.empty(); + } + LocalDateTime processingStartedAt = analysis.startProcessing(now()); + return Optional.of(new AnalysisProcessingClaim(requestJson.get(), processingStartedAt)); + } + + @Transactional + public Optional restartStaleProcessing(Long analysisId, LocalDateTime cutoff) { + Analysis analysis = analysisRepository.findByIdForUpdate(analysisId) + .orElseThrow(() -> new GeneralException(AnalysisErrorStatus.ANALYSIS_NOT_FOUND)); + if (analysis.getStatus() != AnalysisStatus.PROCESSING + || analysis.getProcessingStartedAt() == null + || analysis.getProcessingStartedAt().isAfter(cutoff)) { + return Optional.empty(); + } + Optional requestJson = requestJsonOrFail(analysis); + if (requestJson.isEmpty()) { + return Optional.empty(); + } + LocalDateTime processingStartedAt = analysis.restartProcessing(now()); + return Optional.of(new AnalysisProcessingClaim(requestJson.get(), processingStartedAt)); + } + + @Transactional + public boolean complete( + Long analysisId, + LocalDateTime expectedProcessingStartedAt, + JsonNode result, + String rawResultJson, + GeneratedAnalysisReport generatedReport + ) { + Analysis analysis = getAnalysisForUpdate(analysisId); + if (!analysis.isCurrentProcessing(expectedProcessingStartedAt)) { + return false; + } + validateResult(result); + JsonNode scores = result.path("scores"); + BigDecimal finalScore = requiredScore(scores, "final_score"); + JsonNode domains = scores.path("domains"); + analysis.complete( + finalScore.setScale(0, RoundingMode.HALF_UP).intValueExact(), + resolveGrade(scores.path("grade").asText(), finalScore), + result.path("summary").asText(null), + requiredScore(domains, SCALE), + requiredScore(domains, TENSION), + requiredScore(domains, PROGRESSION), + requiredScore(domains, VOICE_LEADING), + rawResultJson, + now() + ); + AnalysisReport report = generatedReport.generationType() == ReportGenerationType.LLM + ? AnalysisReport.createLlmReport( + analysis, + generatedReport.content(), + generatedReport.modelName(), + generatedReport.promptVersion() + ) + : AnalysisReport.createRuleBasedReport(analysis, generatedReport.content()); + analysisReportRepository.save(report); + saveLlmCallLog(analysis, report, generatedReport.llmCall()); + eventPublisher.publishEvent( + AnalysisCompletedEvent.of(analysis.getUser().getUserId()) + ); + return true; + } + + public void validateResult(JsonNode result) { + if (result == null || !result.isObject()) { + throw invalidRawResult(); + } + JsonNode scores = result.path("scores"); + if (!scores.isObject()) { + throw invalidRawResult(); + } + requiredScore(scores, "final_score"); + JsonNode domains = scores.path("domains"); + if (!domains.isObject()) { + throw invalidRawResult(); + } + requiredScore(domains, SCALE); + requiredScore(domains, TENSION); + requiredScore(domains, PROGRESSION); + requiredScore(domains, VOICE_LEADING); + } + + private Optional requestJsonOrFail(Analysis analysis) { + String requestJson = analysis.getAnalysisRequestJson(); + if (requestJson == null || requestJson.isBlank()) { + analysis.fail(AnalysisErrorStatus.INVALID_ANALYSIS_REQUEST.getMessage(), now()); + return Optional.empty(); + } + return Optional.of(requestJson); + } + + private void saveLlmCallLog( + Analysis analysis, + AnalysisReport report, + LlmCallMetadata metadata + ) { + LlmCallLog log = switch (metadata.status()) { + case SUCCESS -> LlmCallLog.success( + analysis.getUser(), analysis, report, null, + LlmPurpose.REPORT_GENERATION, metadata.modelName(), metadata.promptVersion(), + metadata.promptSnapshot(), metadata.promptTokens(), metadata.completionTokens(), + metadata.totalTokens(), metadata.temperature(), metadata.latencyMs(), + metadata.cacheHit(), metadata.inputHash() + ); + case TIMEOUT -> LlmCallLog.timeout( + analysis.getUser(), analysis, report, null, + LlmPurpose.REPORT_GENERATION, metadata.modelName(), metadata.promptVersion(), + metadata.promptSnapshot(), metadata.temperature(), metadata.latencyMs(), + metadata.inputHash(), metadata.errorMessage() + ); + case FAILED -> LlmCallLog.failed( + analysis.getUser(), analysis, report, null, + LlmPurpose.REPORT_GENERATION, metadata.modelName(), metadata.promptVersion(), + metadata.promptSnapshot(), metadata.temperature(), metadata.latencyMs(), + metadata.inputHash(), metadata.errorMessage() + ); + }; + llmCallLogRepository.save(log); + } + + @Transactional + public boolean fail(Long analysisId, LocalDateTime expectedProcessingStartedAt, String reason) { + Analysis analysis = getAnalysisForUpdate(analysisId); + if (!analysis.isCurrentProcessing(expectedProcessingStartedAt)) { + return false; + } + analysis.fail(reason, now()); + return true; + } + + private LocalDateTime now() { + return LocalDateTime.now(clock); + } + + private Analysis getAnalysisForUpdate(Long analysisId) { + return analysisRepository.findByIdForUpdate(analysisId) + .orElseThrow(() -> new GeneralException(AnalysisErrorStatus.ANALYSIS_NOT_FOUND)); + } + + private BigDecimal requiredScore(JsonNode parent, String fieldName) { + JsonNode node = parent.path(fieldName); + if (!node.isNumber()) { + throw invalidRawResult(); + } + BigDecimal score = node.decimalValue(); + if (score.compareTo(MIN_SCORE) < 0 || score.compareTo(MAX_SCORE) > 0) { + throw invalidRawResult(); + } + return score; + } + + private GeneralException invalidRawResult() { + return new GeneralException(AnalysisErrorStatus.INVALID_RAW_RESULT); + } + + private AnalysisGrade resolveGrade(String value, BigDecimal score) { + return switch (value) { + case "훌륭함", "EXCELLENT" -> AnalysisGrade.EXCELLENT; + case "좋음", "GOOD" -> AnalysisGrade.GOOD; + case "보통", "FAIR" -> AnalysisGrade.FAIR; + case "연습 필요", "POOR" -> AnalysisGrade.POOR; + default -> { + if (score.compareTo(new BigDecimal("90")) >= 0) yield AnalysisGrade.EXCELLENT; + if (score.compareTo(new BigDecimal("75")) >= 0) yield AnalysisGrade.GOOD; + if (score.compareTo(new BigDecimal("60")) >= 0) yield AnalysisGrade.FAIR; + yield AnalysisGrade.POOR; + } + }; + } +} diff --git a/src/main/java/com/mr/domain/analysis/service/ReportGenerationService.java b/src/main/java/com/mr/domain/analysis/service/ReportGenerationService.java new file mode 100644 index 00000000..ad09267d --- /dev/null +++ b/src/main/java/com/mr/domain/analysis/service/ReportGenerationService.java @@ -0,0 +1,174 @@ +package com.mr.domain.analysis.service; + +import com.fasterxml.jackson.databind.JsonNode; +import com.fasterxml.jackson.databind.ObjectMapper; +import com.mr.domain.analysis.entity.enums.ReportGenerationType; +import com.mr.domain.analysis.generator.RuleBasedReportGenerator; +import com.mr.domain.analysis.model.GeneratedAnalysisReport; +import com.mr.domain.analysis.model.LlmCallMetadata; +import com.mr.domain.mentor.entity.enums.LlmCallStatus; +import com.mr.global.client.gemini.GeminiClient; +import com.mr.global.client.gemini.GeminiGenerationResult; +import com.mr.global.config.GeminiProperties; +import java.math.BigDecimal; +import java.net.SocketTimeoutException; +import java.net.http.HttpTimeoutException; +import java.nio.charset.StandardCharsets; +import java.security.MessageDigest; +import java.security.NoSuchAlgorithmException; +import java.util.List; +import lombok.RequiredArgsConstructor; +import lombok.extern.slf4j.Slf4j; +import org.springframework.stereotype.Service; + +@Slf4j +@Service +@RequiredArgsConstructor +public class ReportGenerationService { + + static final String PROMPT_VERSION = "analysis-report-v2"; + private static final BigDecimal TEMPERATURE = new BigDecimal("0.30"); + private static final int MIN_REPORT_LENGTH = 600; + private static final List REQUIRED_HEADINGS = List.of( + "# 연주 분석 리포트", + "## 총평", + "## 잘한 점", + "## 진행 맥락", + "## 개선 제안", + "## 점수 요약" + ); + private static final String SYSTEM_PROMPT = """ + 당신은 재즈 화성학과 MIDI 연주 분석에 능숙한 친절한 음악 코치입니다. + 입력은 MuseReview 분석 서버가 생성한 JSON입니다. JSON에 존재하는 사실만 사용하세요. + 응답은 한국어 Markdown으로 작성하고 반드시 다음 순서를 지키세요. + + # 연주 분석 리포트 + **조성** ... · **장르** ... · **박자** ... · **템포** ... bpm + ## 총평 + ## 잘한 점 + ## 진행 맥락 + ## 개선 제안 + ## 점수 요약 + + 섹션별 작성 기준은 다음과 같습니다. + - 총평: 입력의 summary와 종합 점수를 바탕으로 3~4문장, 최소 150자 + - 잘한 점: 근거가 있는 범위에서 2~3개 항목, 각 항목은 최소 100자로 영역 점수나 구체적 분석 근거 포함 + - 진행 맥락: 근거가 있는 범위에서 2~3개 항목, 각 항목은 최소 100자로 마디·코드 진행·음표 중 입력에 존재하는 근거 포함 + - 개선 제안: 2~3개 항목, 각 항목은 최소 100자로 문제점·근거·실행 가능한 연습 방법 포함 + - 점수 요약: 종합 점수와 네 영역 점수를 입력값 그대로 표시 + - 전체 본문: 700자 이상 1,500자 이하 + + 사용자를 비난하지 말고 구체적인 마디·코드·음표 근거를 우선 제시하세요. + 점수와 수치를 임의로 만들거나 변경하지 마세요. + 근거가 부족하면 항목 수를 억지로 채우지 말고, 입력 JSON에 없는 사실을 만들지 마세요. + """; + + private final GeminiClient geminiClient; + private final GeminiProperties properties; + private final RuleBasedReportGenerator ruleBasedReportGenerator; + private final ObjectMapper objectMapper; + + public GeneratedAnalysisReport generate(JsonNode analysisResult) { + long startedAt = System.nanoTime(); + String input = analysisResult.toString(); + JsonNode promptSnapshot = promptSnapshot(analysisResult); + String inputHash = sha256(PROMPT_VERSION + ":" + input); + try { + GeminiGenerationResult result = geminiClient.generateReport(SYSTEM_PROMPT, input); + validateMarkdownStructure(result.content()); + return new GeneratedAnalysisReport( + ReportGenerationType.LLM, + result.content(), + properties.model(), + PROMPT_VERSION, + new LlmCallMetadata( + LlmCallStatus.SUCCESS, + properties.model(), + PROMPT_VERSION, + promptSnapshot, + result.promptTokens(), + result.completionTokens(), + result.totalTokens(), + TEMPERATURE, + elapsedMillis(startedAt), + result.cacheHit(), + inputHash, + null + ) + ); + } catch (Exception exception) { + log.warn("Gemini report generation failed; using rule-based fallback.", exception); + return new GeneratedAnalysisReport( + ReportGenerationType.RULE_BASED, + ruleBasedReportGenerator.generate(analysisResult), + null, + PROMPT_VERSION, + new LlmCallMetadata( + isTimeout(exception) ? LlmCallStatus.TIMEOUT : LlmCallStatus.FAILED, + properties.model(), + PROMPT_VERSION, + promptSnapshot, + null, + null, + null, + TEMPERATURE, + elapsedMillis(startedAt), + false, + inputHash, + exception.getMessage() + ) + ); + } + } + + private void validateMarkdownStructure(String content) { + if (content == null || content.strip().length() < MIN_REPORT_LENGTH) { + throw new IllegalStateException("Gemini returned a report that is too short."); + } + List lines = content.lines() + .map(String::stripTrailing) + .toList(); + int previousIndex = -1; + for (String heading : REQUIRED_HEADINGS) { + int currentIndex = lines.subList(previousIndex + 1, lines.size()).indexOf(heading); + if (currentIndex < 0) { + throw new IllegalStateException("Gemini returned an invalid report structure."); + } + previousIndex += currentIndex + 1; + } + } + + private JsonNode promptSnapshot(JsonNode analysisResult) { + var snapshot = objectMapper.createObjectNode(); + snapshot.put("promptVersion", PROMPT_VERSION); + snapshot.put("systemPrompt", SYSTEM_PROMPT); + snapshot.set("analysisResult", analysisResult); + return snapshot; + } + + private int elapsedMillis(long startedAt) { + long elapsed = (System.nanoTime() - startedAt) / 1_000_000L; + return (int) Math.min(elapsed, Integer.MAX_VALUE); + } + + private boolean isTimeout(Throwable throwable) { + Throwable current = throwable; + while (current != null) { + if (current instanceof SocketTimeoutException || current instanceof HttpTimeoutException) { + return true; + } + current = current.getCause(); + } + return false; + } + + private String sha256(String value) { + try { + byte[] digest = MessageDigest.getInstance("SHA-256") + .digest(value.getBytes(StandardCharsets.UTF_8)); + return java.util.HexFormat.of().formatHex(digest); + } catch (NoSuchAlgorithmException exception) { + throw new IllegalStateException("SHA-256 is not available.", exception); + } + } +} diff --git a/src/main/java/com/mr/domain/mentor/repository/LlmCallLogRepository.java b/src/main/java/com/mr/domain/mentor/repository/LlmCallLogRepository.java new file mode 100644 index 00000000..759e4249 --- /dev/null +++ b/src/main/java/com/mr/domain/mentor/repository/LlmCallLogRepository.java @@ -0,0 +1,7 @@ +package com.mr.domain.mentor.repository; + +import com.mr.domain.mentor.entity.LlmCallLog; +import org.springframework.data.jpa.repository.JpaRepository; + +public interface LlmCallLogRepository extends JpaRepository { +} diff --git a/src/main/java/com/mr/domain/playing/exception/PlayingErrorStatus.java b/src/main/java/com/mr/domain/playing/exception/PlayingErrorStatus.java index c5379fb1..7fc01502 100644 --- a/src/main/java/com/mr/domain/playing/exception/PlayingErrorStatus.java +++ b/src/main/java/com/mr/domain/playing/exception/PlayingErrorStatus.java @@ -14,12 +14,11 @@ public enum PlayingErrorStatus implements BaseCode { MISSING_BACKING_TRACK_ID(HttpStatus.BAD_REQUEST, "PLAYING_400_03", "백킹트랙 연주 모드에서는 백킹트랙 ID가 필수입니다."), MISSING_PLAYING(HttpStatus.BAD_REQUEST, "PLAYING_400_04", "MIDI 이벤트를 기록할 연주 정보가 누락되었습니다."), MISSING_MIDI_TYPE(HttpStatus.BAD_REQUEST, "PLAYING_400_05", "MIDI Type은 필수 입력 값입니다."), - // 기존 PLAYING_400_06 ~ PLAYING_400_10은 MidiEventErrorStatus로 분리 MISSING_PLAYING_MODE(HttpStatus.BAD_REQUEST, "PLAYING_400_11", "연주 모드는 필수 입력값입니다."), MISSING_PLAYING_STATUS(HttpStatus.BAD_REQUEST, "PLAYING_400_12", "연주 상태는 필수 입력값입니다."), UNSUPPORTED_PLAYING_MODE(HttpStatus.BAD_REQUEST, "PLAYING_400_13", "현재 지원하지 않는 연주 모드입니다."), - PLAYING_ACCESS_DENIED(HttpStatus.FORBIDDEN, "PLAYING_403_01", "해당 연주에 대한 접근 권한이 없습니다."), - PLAYING_NOT_FOUND(HttpStatus.NOT_FOUND, "PLAYING_404_01", "연주 세션을 찾을 수 없습니다."), + PLAYING_ACCESS_DENIED(HttpStatus.FORBIDDEN, "PLAYING_403_01", "해당 연주 기록에 접근할 수 없습니다."), + PLAYING_NOT_FOUND(HttpStatus.NOT_FOUND, "PLAYING_404_01", "연주 기록을 찾을 수 없습니다."), INVALID_PLAYING_STATUS(HttpStatus.CONFLICT, "PLAYING_409_01", "현재 연주 상태에서는 요청한 작업을 수행할 수 없습니다."), MISSING_PLAYING_START_TIME(HttpStatus.CONFLICT, "PLAYING_409_02", "연주 시작 시간이 기록되지 않았습니다."), INVALID_PLAYING_DURATION(HttpStatus.CONFLICT, "PLAYING_409_03", "연주 종료 시간이 시작 시간보다 이전일 수 없습니다."), diff --git a/src/main/java/com/mr/domain/playing/repository/PlayingRepository.java b/src/main/java/com/mr/domain/playing/repository/PlayingRepository.java index 590d44b1..a7188404 100644 --- a/src/main/java/com/mr/domain/playing/repository/PlayingRepository.java +++ b/src/main/java/com/mr/domain/playing/repository/PlayingRepository.java @@ -2,6 +2,7 @@ import com.mr.domain.playing.entity.Playing; import com.mr.domain.playing.entity.enums.PlayingStatus; +import jakarta.persistence.LockModeType; import java.time.LocalDate; import java.time.LocalDateTime; import java.util.List; @@ -9,6 +10,7 @@ import org.springframework.data.domain.Pageable; import org.springframework.data.domain.Slice; import org.springframework.data.jpa.repository.JpaRepository; +import org.springframework.data.jpa.repository.Lock; import org.springframework.data.jpa.repository.Query; import org.springframework.data.repository.query.Param; @@ -38,6 +40,15 @@ Slice findPlayingsByUserAndStatus( """) Optional findByIdWithBackingTrack(@Param("id") Long id); + @Lock(LockModeType.PESSIMISTIC_WRITE) + @Query(""" + select p from Playing p + left join fetch p.backingTrack + where p.id = :id + and p.deletedAt is null + """) + Optional findByIdWithBackingTrackForUpdate(@Param("id") Long id); + @Query(""" select p from Playing p where p.user.userId = :userId @@ -52,7 +63,6 @@ List findByUserAndStatusSince( @Param("since") LocalDateTime since ); - // 연속 출석일수는 상한이 없어 기간 제한 없이 날짜 단위 distinct 조회 @Query(""" select distinct function('date', p.endedAt) from Playing p where p.user.userId = :userId @@ -65,7 +75,6 @@ List findDistinctEndedDatesByUserAndStatus( @Param("status") PlayingStatus status ); - // 통계 집계용 전체 기간 요약 - row를 끌어오지 않고 단일 행으로 집계 @Query(""" select count(p) as sessionCount, coalesce(sum(p.durationSec), 0) as totalDurationSec, diff --git a/src/main/java/com/mr/domain/statistics/entity/PracticeStatistics.java b/src/main/java/com/mr/domain/statistics/entity/PracticeStatistics.java index 7faca1d5..2a599371 100644 --- a/src/main/java/com/mr/domain/statistics/entity/PracticeStatistics.java +++ b/src/main/java/com/mr/domain/statistics/entity/PracticeStatistics.java @@ -60,11 +60,11 @@ public class PracticeStatistics extends BaseCreatedEntity { @Column(name = "session_count", nullable = false) private Integer sessionCount; - // 산출 기준 미확정(PM 확인 필요) - 집계 로직에서 채우지 않음 + // 미집계 지표 @Column(name = "average_score", precision = 5, scale = 2) private BigDecimal averageScore; - // COMPLETED Analysis의 totalScore 평균 (해당 기간) + // 기간별 완료 분석 평균 점수 @Column(name = "average_accuracy", precision = 5, scale = 2) private BigDecimal averageAccuracy; diff --git a/src/main/java/com/mr/domain/statistics/entity/UserStatistics.java b/src/main/java/com/mr/domain/statistics/entity/UserStatistics.java index 5531302e..efae99e2 100644 --- a/src/main/java/com/mr/domain/statistics/entity/UserStatistics.java +++ b/src/main/java/com/mr/domain/statistics/entity/UserStatistics.java @@ -43,11 +43,11 @@ public class UserStatistics extends BaseTimeEntity { @Column(name = "total_analysis_count", nullable = false) private Integer totalAnalysisCount; - // 산출 기준 미확정(PM 확인 필요) - 집계 로직에서 채우지 않음 + // 미집계 지표 @Column(name = "average_score", precision = 5, scale = 2) private BigDecimal averageScore; - // COMPLETED Analysis의 totalScore 평균 (전체 기간) + // 전체 완료 분석 평균 점수 @Column(name = "average_accuracy", precision = 5, scale = 2) private BigDecimal averageAccuracy; diff --git a/src/main/java/com/mr/domain/statistics/service/AnalysisSkillScoreResolver.java b/src/main/java/com/mr/domain/statistics/service/AnalysisSkillScoreResolver.java index dd646c8b..e8a8837b 100644 --- a/src/main/java/com/mr/domain/statistics/service/AnalysisSkillScoreResolver.java +++ b/src/main/java/com/mr/domain/statistics/service/AnalysisSkillScoreResolver.java @@ -4,7 +4,6 @@ import com.mr.domain.statistics.entity.enums.SkillType; import java.math.BigDecimal; -// StatisticsService(라이브 조회)와 StatisticsAggregationService(집계 write)가 공유하는 스킬별 점수 추출 final class AnalysisSkillScoreResolver { private AnalysisSkillScoreResolver() { diff --git a/src/main/java/com/mr/domain/statistics/service/StatisticsAggregationService.java b/src/main/java/com/mr/domain/statistics/service/StatisticsAggregationService.java index a6f87f91..cc007c25 100644 --- a/src/main/java/com/mr/domain/statistics/service/StatisticsAggregationService.java +++ b/src/main/java/com/mr/domain/statistics/service/StatisticsAggregationService.java @@ -30,7 +30,6 @@ import org.springframework.transaction.annotation.Propagation; import org.springframework.transaction.annotation.Transactional; -// 연습/분석 완료 이벤트를 받아 UserStatistics/PracticeStatistics/SkillStatistics를 다시 계산해 저장 @Service @RequiredArgsConstructor @Transactional @@ -47,8 +46,7 @@ public class StatisticsAggregationService { private final AnalysisRepository analysisRepository; private final Clock clock; - // 리스너의 @Retryable과 같은 메서드에 두면 advice 순서가 불명확해지므로, - // 트랜잭션 경계는 반드시 이 별도 빈의 메서드에서 열어 재시도마다 새 트랜잭션이 보장되도록 함 + // 재시도별 신규 트랜잭션 보장 @Transactional(propagation = Propagation.REQUIRES_NEW) public void onPlayingCompleted(Long userId) { refreshUserPracticeTotals(userId); @@ -83,7 +81,6 @@ private void refreshUserAnalysisTotals(Long userId) { toScoreScale(totals.getAverageTotalScore())); } - // 동시 이벤트로 인해 유니크 제약을 어겨 저장에 실패해도, 리스너의 재시도가 재조회 후 갱신으로 이어져 정상 처리됨 private UserStatistics getOrCreateUserStatistics(Long userId) { return userStatisticsRepository.findByUser_UserId(userId) .orElseGet(() -> userStatisticsRepository.save(UserStatistics.createForUser(getUser(userId)))); @@ -127,7 +124,7 @@ private void upsertWeeklySkillStatistics(Long userId, SkillType skillType, Local LocalDate weekEnd, LocalDate lastWeekStart, List analyses) { BigDecimal score = averageSkillScore(analyses, skillType); if (score == null) { - // score는 NOT NULL 컬럼 - 이번 주 유효한 점수가 없으면 생성/갱신하지 않고 건너뜀 + // NOT NULL 제약에 따른 미집계 처리 return; } diff --git a/src/main/java/com/mr/domain/statistics/service/StatisticsEventListener.java b/src/main/java/com/mr/domain/statistics/service/StatisticsEventListener.java index e1c7cb4d..e66ea984 100644 --- a/src/main/java/com/mr/domain/statistics/service/StatisticsEventListener.java +++ b/src/main/java/com/mr/domain/statistics/service/StatisticsEventListener.java @@ -13,7 +13,6 @@ import org.springframework.transaction.event.TransactionPhase; import org.springframework.transaction.event.TransactionalEventListener; -// 트랜잭션 경계는 StatisticsAggregationService에서 관리(별도 빈으로 분리한 이유는 그쪽 주석 참고) @Slf4j @Component @RequiredArgsConstructor @@ -21,7 +20,7 @@ public class StatisticsEventListener { private final StatisticsAggregationService statisticsAggregationService; - // 유니크 제약 충돌(동시 upsert)만 재시도 대상 - 그 외 예외(유저 없음 등)는 즉시 전파해 조용히 삼키지 않음 + // 동시 upsert 충돌 한정 재시도 @Retryable( retryFor = {DataIntegrityViolationException.class}, maxAttempts = 3, diff --git a/src/main/java/com/mr/domain/statistics/service/StatisticsService.java b/src/main/java/com/mr/domain/statistics/service/StatisticsService.java index f46a768c..6f3dcc88 100644 --- a/src/main/java/com/mr/domain/statistics/service/StatisticsService.java +++ b/src/main/java/com/mr/domain/statistics/service/StatisticsService.java @@ -65,7 +65,6 @@ public StatisticsResponseDTO getStatistics(Long userId) { ); } - // index 0=이번 주, 1=지난 주, 2=2주 전, 3=3주 전 private ScoreAggregate[] buildWeeklyScoreAggregates(List analyses, LocalDateTime thisWeekStart) { ScoreAggregate[] aggregates = new ScoreAggregate[WEEKLY_TREND_WEEKS]; for (int weeksAgo = 0; weeksAgo < WEEKLY_TREND_WEEKS; weeksAgo++) { @@ -201,7 +200,6 @@ private int diffIntOrZero(ScoreAggregate current, ScoreAggregate previous) { .intValue(); } - // 평균 지표의 분모(count) 0 처리 공통화 private record ScoreAggregate(BigDecimal sum, int count) { BigDecimal average() { if (count == 0) { diff --git a/src/main/java/com/mr/global/client/ai/AiAnalysisRequest.java b/src/main/java/com/mr/global/client/ai/AiAnalysisRequest.java index 8388f3c3..63fb804d 100644 --- a/src/main/java/com/mr/global/client/ai/AiAnalysisRequest.java +++ b/src/main/java/com/mr/global/client/ai/AiAnalysisRequest.java @@ -13,7 +13,8 @@ public record Meta( Double bpm, @JsonProperty("time_signature") List timeSignature, Key key, - String genre + String genre, + String level ) { } @@ -32,10 +33,15 @@ public record Chord( public record Note( Integer index, + NoteType type, Integer pitch, - @JsonProperty("onset_beats") Double onsetBeats, - @JsonProperty("duration_beats") Double durationBeats, - Integer velocity + Integer velocity, + @JsonProperty("timestamp_ms") Double timestampMs ) { } + + public enum NoteType { + NOTE_ON, + NOTE_OFF + } } diff --git a/src/main/java/com/mr/global/client/gemini/GeminiClient.java b/src/main/java/com/mr/global/client/gemini/GeminiClient.java new file mode 100644 index 00000000..91cefaea --- /dev/null +++ b/src/main/java/com/mr/global/client/gemini/GeminiClient.java @@ -0,0 +1,80 @@ +package com.mr.global.client.gemini; + +import com.fasterxml.jackson.databind.JsonNode; +import com.mr.global.config.GeminiProperties; +import java.util.List; +import java.util.Map; +import lombok.RequiredArgsConstructor; +import org.springframework.stereotype.Component; +import org.springframework.web.client.RestClient; + +@Component +@RequiredArgsConstructor +public class GeminiClient { + + private final RestClient geminiRestClient; + private final GeminiProperties properties; + + public GeminiGenerationResult generateReport(String systemPrompt, String analysisJson) { + if (properties.apiKey() == null || properties.apiKey().isBlank()) { + throw new IllegalStateException("Gemini API key is not configured."); + } + + Map body = Map.of( + "system_instruction", content(systemPrompt), + "contents", List.of(Map.of( + "role", "user", + "parts", List.of(Map.of("text", analysisJson)) + )), + "generationConfig", Map.of( + "temperature", 0.3, + "maxOutputTokens", 4_096 + ) + ); + + JsonNode response = geminiRestClient.post() + .uri("/v1beta/models/{model}:generateContent", properties.model()) + .header("x-goog-api-key", properties.apiKey()) + .body(body) + .retrieve() + .body(JsonNode.class); + + String content = extractText(response); + if (content == null || content.isBlank()) { + throw new IllegalStateException("Gemini returned an empty report."); + } + JsonNode usage = response.path("usageMetadata"); + return new GeminiGenerationResult( + content.trim(), + integerOrNull(usage.path("promptTokenCount")), + integerOrNull(usage.path("candidatesTokenCount")), + integerOrNull(usage.path("totalTokenCount")), + usage.path("cachedContentTokenCount").asInt(0) > 0 + ); + } + + private Map content(String text) { + return Map.of("parts", List.of(Map.of("text", text))); + } + + private String extractText(JsonNode response) { + if (response == null) { + return null; + } + JsonNode parts = response.path("candidates").path(0).path("content").path("parts"); + if (!parts.isArray()) { + return null; + } + StringBuilder text = new StringBuilder(); + for (JsonNode part : parts) { + if (part.path("text").isTextual()) { + text.append(part.path("text").asText()); + } + } + return text.toString(); + } + + private Integer integerOrNull(JsonNode node) { + return node.canConvertToInt() ? node.intValue() : null; + } +} diff --git a/src/main/java/com/mr/global/client/gemini/GeminiGenerationResult.java b/src/main/java/com/mr/global/client/gemini/GeminiGenerationResult.java new file mode 100644 index 00000000..914d1d1e --- /dev/null +++ b/src/main/java/com/mr/global/client/gemini/GeminiGenerationResult.java @@ -0,0 +1,10 @@ +package com.mr.global.client.gemini; + +public record GeminiGenerationResult( + String content, + Integer promptTokens, + Integer completionTokens, + Integer totalTokens, + boolean cacheHit +) { +} diff --git a/src/main/java/com/mr/global/config/GeminiProperties.java b/src/main/java/com/mr/global/config/GeminiProperties.java new file mode 100644 index 00000000..ca0b22c4 --- /dev/null +++ b/src/main/java/com/mr/global/config/GeminiProperties.java @@ -0,0 +1,24 @@ +package com.mr.global.config; + +import java.time.Duration; +import org.springframework.boot.context.properties.ConfigurationProperties; + +@ConfigurationProperties(prefix = "external.ai") +public record GeminiProperties( + String baseUrl, + String apiKey, + String model, + Duration connectTimeout, + Duration readTimeout +) { + public GeminiProperties { + baseUrl = hasText(baseUrl) ? baseUrl : "https://generativelanguage.googleapis.com"; + model = hasText(model) ? model : "gemini-3-flash-preview"; + connectTimeout = connectTimeout != null ? connectTimeout : Duration.ofSeconds(5); + readTimeout = readTimeout != null ? readTimeout : Duration.ofSeconds(60); + } + + private static boolean hasText(String value) { + return value != null && !value.isBlank(); + } +} diff --git a/src/main/java/com/mr/global/config/GeminiRestClientConfig.java b/src/main/java/com/mr/global/config/GeminiRestClientConfig.java new file mode 100644 index 00000000..cb52552f --- /dev/null +++ b/src/main/java/com/mr/global/config/GeminiRestClientConfig.java @@ -0,0 +1,26 @@ +package com.mr.global.config; + +import org.springframework.boot.context.properties.EnableConfigurationProperties; +import org.springframework.boot.web.client.ClientHttpRequestFactories; +import org.springframework.boot.web.client.ClientHttpRequestFactorySettings; +import org.springframework.context.annotation.Bean; +import org.springframework.context.annotation.Configuration; +import org.springframework.web.client.RestClient; + +@Configuration +@EnableConfigurationProperties(GeminiProperties.class) +public class GeminiRestClientConfig { + + @Bean + public RestClient geminiRestClient(GeminiProperties properties) { + ClientHttpRequestFactorySettings settings = ClientHttpRequestFactorySettings.DEFAULTS + .withConnectTimeout(properties.connectTimeout()) + .withReadTimeout(properties.readTimeout()); + + return RestClient.builder() + .baseUrl(properties.baseUrl()) + .requestFactory(ClientHttpRequestFactories.get(settings)) + .defaultHeader("Content-Type", "application/json") + .build(); + } +} diff --git a/src/main/java/com/mr/global/event/AnalysisCompletedEvent.java b/src/main/java/com/mr/global/event/AnalysisCompletedEvent.java index 37f4588d..517ca946 100644 --- a/src/main/java/com/mr/global/event/AnalysisCompletedEvent.java +++ b/src/main/java/com/mr/global/event/AnalysisCompletedEvent.java @@ -9,7 +9,7 @@ public class AnalysisCompletedEvent { private Long userId; - // userId는 완료 처리된 Analysis의 소유자(analysis.getUser().getUserId())에서만 채울 것 - 클라이언트 입력 금지 + // 완료 분석 소유자 식별자 public static AnalysisCompletedEvent of(Long userId) { return new AnalysisCompletedEvent(userId); } diff --git a/src/main/java/com/mr/global/event/PlayingCompletedEvent.java b/src/main/java/com/mr/global/event/PlayingCompletedEvent.java index ba1484ea..45e81259 100644 --- a/src/main/java/com/mr/global/event/PlayingCompletedEvent.java +++ b/src/main/java/com/mr/global/event/PlayingCompletedEvent.java @@ -9,7 +9,7 @@ public class PlayingCompletedEvent { private Long userId; - // userId는 완료 처리된 Playing의 소유자(playing.getUser().getUserId())에서만 채울 것 - 클라이언트 입력 금지 + // 클라이언트 입력이 아닌 완료된 Playing 소유자 식별자 public static PlayingCompletedEvent of(Long userId) { return new PlayingCompletedEvent(userId); } diff --git a/src/test/java/com/mr/domain/analysis/controller/AnalysisControllerTest.java b/src/test/java/com/mr/domain/analysis/controller/AnalysisControllerTest.java index 217770e5..16dc1be8 100644 --- a/src/test/java/com/mr/domain/analysis/controller/AnalysisControllerTest.java +++ b/src/test/java/com/mr/domain/analysis/controller/AnalysisControllerTest.java @@ -117,6 +117,12 @@ void getAnalysisResult_success() throws Exception { AnalysisResultResponseDTO response = new AnalysisResultResponseDTO( 1L, 1L, + "Jazz Standard Practice", + "jazz", + "C Major", + 120, + LocalDateTime.of(2026, 7, 24, 9, 0), + AnalysisStatus.COMPLETED, 1, 8, 85, @@ -149,9 +155,13 @@ void getAnalysisResult_success() throws Exception { .andExpect(status().isOk()) .andExpect(jsonPath("$.isSuccess").value(true)) .andExpect(jsonPath("$.data.analysisId").value(1L)) + .andExpect(jsonPath("$.data.title").value("Jazz Standard Practice")) + .andExpect(jsonPath("$.data.status").value("COMPLETED")) .andExpect(jsonPath("$.data.totalScore").value(85)) .andExpect(jsonPath("$.data.grade").value("GOOD")) - .andExpect(jsonPath("$.data.domainScores.scaleScore").value(80.00)) + .andExpect(jsonPath("$.data.domainScores.scale").value(80.00)) + .andExpect(jsonPath("$.data.domainScores.scaleScore").doesNotExist()) + .andExpect(jsonPath("$.data.result").doesNotExist()) .andExpect(jsonPath("$.data.report.modelName").value("gpt-4")); } @@ -163,6 +173,6 @@ void getAnalysisResult_notCompleted() throws Exception { mockMvc.perform(get("/api/analyses/{analysisId}", 1L)) .andExpect(status().isConflict()) - .andExpect(jsonPath("$.code").value("ANALYSIS_409_01")); + .andExpect(jsonPath("$.code").value("ANALYSIS_409_02")); } } diff --git a/src/test/java/com/mr/domain/analysis/entity/AnalysisTest.java b/src/test/java/com/mr/domain/analysis/entity/AnalysisTest.java index 83398a34..d76685ff 100644 --- a/src/test/java/com/mr/domain/analysis/entity/AnalysisTest.java +++ b/src/test/java/com/mr/domain/analysis/entity/AnalysisTest.java @@ -10,11 +10,14 @@ import com.mr.domain.playing.entity.Playing; import com.mr.domain.user.entity.User; import com.mr.global.apipayload.exception.GeneralException; +import java.time.LocalDateTime; import org.junit.jupiter.api.DisplayName; import org.junit.jupiter.api.Test; class AnalysisTest { + private static final LocalDateTime NOW = LocalDateTime.of(2026, 7, 31, 12, 0); + @Test @DisplayName("createPending - user가 null이면 예외가 발생한다") void createPending_userNull_throwsException() { @@ -58,6 +61,36 @@ void createPending_success_setsPendingStatus() { assertThat(analysis.getStatus()).isEqualTo(AnalysisStatus.PENDING); } + @Test + @DisplayName("startProcessing - 처리 시작 시각을 기록한다") + void startProcessing_recordsProcessingStartedAt() { + User user = mock(User.class); + Playing playing = mock(Playing.class); + given(playing.getUser()).willReturn(user); + Analysis analysis = Analysis.createPending(user, playing, 1, 8, "{}"); + + analysis.startProcessing(NOW); + + assertThat(analysis.getProcessingStartedAt()).isEqualTo(NOW); + assertThat(analysis.getStatus()).isEqualTo(AnalysisStatus.PROCESSING); + } + + @Test + @DisplayName("restartProcessing - 새 처리 시작 시각으로 이전 작업을 펜싱한다") + void restartProcessing_renewsProcessingStartedAt() { + User user = mock(User.class); + Playing playing = mock(Playing.class); + given(playing.getUser()).willReturn(user); + Analysis analysis = Analysis.createPending(user, playing, 1, 8, "{}"); + LocalDateTime firstAttempt = analysis.startProcessing(NOW); + + LocalDateTime secondAttempt = analysis.restartProcessing(NOW); + + assertThat(secondAttempt).isAfter(firstAttempt); + assertThat(analysis.isCurrentProcessing(firstAttempt)).isFalse(); + assertThat(analysis.isCurrentProcessing(secondAttempt)).isTrue(); + } + @Test @DisplayName("createPending - playing 소유자와 user가 다르면 예외가 발생한다") void createPending_ownerMismatch_throwsException() { diff --git a/src/test/java/com/mr/domain/analysis/generator/AnalysisResultEnricherTest.java b/src/test/java/com/mr/domain/analysis/generator/AnalysisResultEnricherTest.java new file mode 100644 index 00000000..c5fe5628 --- /dev/null +++ b/src/test/java/com/mr/domain/analysis/generator/AnalysisResultEnricherTest.java @@ -0,0 +1,79 @@ +package com.mr.domain.analysis.generator; + +import static org.assertj.core.api.Assertions.assertThat; +import static org.assertj.core.api.Assertions.assertThatThrownBy; + +import com.fasterxml.jackson.databind.JsonNode; +import com.fasterxml.jackson.databind.ObjectMapper; +import org.junit.jupiter.api.Test; + +class AnalysisResultEnricherTest { + + private final ObjectMapper objectMapper = new ObjectMapper(); + private final AnalysisResultEnricher enricher = new AnalysisResultEnricher(); + + @Test + void enrich_addsEvidenceBasedSummaryWhenMissing() throws Exception { + JsonNode original = resultWithoutSummary(); + + JsonNode enriched = enricher.enrich(original); + + assertThat(enriched.path("summary").asText()) + .isEqualTo("조성에 어울리는 음 선택이 안정적으로 이어졌어요.") + .doesNotContain("80.5", "점", "좋음") + .hasSizeBetween(20, 40); + assertThat(original.has("summary")).isFalse(); + } + + @Test + void enrich_describesTensionWithoutRepeatingScore() throws Exception { + JsonNode original = resultWithoutSummary(); + ((com.fasterxml.jackson.databind.node.ObjectNode) original.path("scores").path("domains")) + .put("텐션", 100); + + JsonNode enriched = enricher.enrich(original); + + assertThat(enriched.path("summary").asText()) + .isEqualTo("긴장감을 살리는 텐션 활용이 자연스럽게 이어졌어요.") + .doesNotContain("100", "점"); + } + + @Test + void enrich_preservesExistingSummary() throws Exception { + JsonNode original = resultWithoutSummary(); + ((com.fasterxml.jackson.databind.node.ObjectNode) original).put("summary", "AI 서버 요약"); + + JsonNode enriched = enricher.enrich(original); + + assertThat(enriched).isSameAs(original); + assertThat(enriched.path("summary").asText()).isEqualTo("AI 서버 요약"); + } + + @Test + void enrich_withoutDomainScores_throwsDescriptiveException() throws Exception { + JsonNode result = objectMapper.readTree(""" + {"scores":{"domains":{}}} + """); + + assertThatThrownBy(() -> enricher.enrich(result)) + .isInstanceOf(IllegalStateException.class) + .hasMessage("도메인 점수가 없습니다."); + } + + private JsonNode resultWithoutSummary() throws Exception { + return objectMapper.readTree(""" + { + "scores": { + "final_score": 80.5, + "grade": "좋음", + "domains": { + "스케일": 95, + "텐션": 60, + "진행": 85, + "코드 연결": 75 + } + } + } + """); + } +} diff --git a/src/test/java/com/mr/domain/analysis/generator/RuleBasedReportGeneratorTest.java b/src/test/java/com/mr/domain/analysis/generator/RuleBasedReportGeneratorTest.java new file mode 100644 index 00000000..070ce249 --- /dev/null +++ b/src/test/java/com/mr/domain/analysis/generator/RuleBasedReportGeneratorTest.java @@ -0,0 +1,46 @@ +package com.mr.domain.analysis.generator; + +import static org.assertj.core.api.Assertions.assertThat; + +import com.fasterxml.jackson.databind.JsonNode; +import com.fasterxml.jackson.databind.ObjectMapper; +import org.junit.jupiter.api.Test; + +class RuleBasedReportGeneratorTest { + + private final ObjectMapper objectMapper = new ObjectMapper(); + private final RuleBasedReportGenerator generator = new RuleBasedReportGenerator(); + + @Test + void generate_usesAvailableAnalysisEvidenceAndRemovesDuplicateRules() throws Exception { + JsonNode result = objectMapper.readTree(""" + { + "meta": {"key":"C major","genre":"jazz","time_signature":[4,4],"bpm":120}, + "summary": "코드 진행은 안정적이지만 코드 연결을 보완하면 좋습니다.", + "scores": { + "final_score": 80, + "domains": {"스케일":90,"텐션":70,"진행":85,"코드 연결":60}, + "coverage": {"note":"충분한 길이로 분석되었습니다."}, + "scale_appropriateness": { + "out_of_scale_notes": [ + {"suggestion":"B음을 Bb로 해결하세요."}, + {"suggestion":"F#음을 G로 해결하세요."} + ] + } + }, + "harmonic_rules": [{"label":"ii-V-I"},{"label":"ii-V-I"}], + "timing_deviations": {"summary": {"flagged_count": 4}}, + "learning_recommendations": [ + {"title":"보이스 리딩","reason":"도약이 큽니다.","study_tip":"가까운 코드톤을 연결하세요."} + ] + } + """); + + String report = generator.generate(result); + + assertThat(report) + .contains("코드 진행은 안정적", "충분한 길이", "스케일 밖 음이 2개") + .contains("박자가 흔들린 음이 4개", "보이스 리딩", "가까운 코드톤") + .containsOnlyOnce("- ii-V-I"); + } +} diff --git a/src/test/java/com/mr/domain/analysis/service/AnalysisProcessingServiceTest.java b/src/test/java/com/mr/domain/analysis/service/AnalysisProcessingServiceTest.java new file mode 100644 index 00000000..82f1c894 --- /dev/null +++ b/src/test/java/com/mr/domain/analysis/service/AnalysisProcessingServiceTest.java @@ -0,0 +1,191 @@ +package com.mr.domain.analysis.service; + +import static org.mockito.BDDMockito.given; +import static org.mockito.Mockito.doThrow; +import static org.mockito.Mockito.never; +import static org.mockito.Mockito.verify; + +import com.fasterxml.jackson.databind.JsonNode; +import com.fasterxml.jackson.databind.ObjectMapper; +import com.mr.domain.analysis.exception.AnalysisErrorStatus; +import com.mr.domain.analysis.entity.enums.ReportGenerationType; +import com.mr.domain.analysis.generator.AnalysisResultEnricher; +import com.mr.domain.analysis.model.AnalysisProcessingClaim; +import com.mr.domain.analysis.model.GeneratedAnalysisReport; +import com.mr.domain.analysis.model.LlmCallMetadata; +import com.mr.domain.mentor.entity.enums.LlmCallStatus; +import com.mr.global.apipayload.exception.GeneralException; +import com.mr.global.client.ai.AiAnalysisRequest; +import com.mr.global.client.ai.AiServerClient; +import java.math.BigDecimal; +import java.time.LocalDateTime; +import java.util.Optional; +import org.junit.jupiter.api.Test; +import org.junit.jupiter.api.extension.ExtendWith; +import org.mockito.Mock; +import org.mockito.junit.jupiter.MockitoExtension; + +@ExtendWith(MockitoExtension.class) +class AnalysisProcessingServiceTest { + + @Mock + private AnalysisStateService analysisStateService; + + @Mock + private AiServerClient aiServerClient; + + @Mock + private ReportGenerationService reportGenerationService; + + @Test + void process_doesNotCallAiWhenAnotherWorkerAlreadyClaimedAnalysis() { + given(analysisStateService.startProcessing(1L)).willReturn(Optional.empty()); + AnalysisProcessingService service = new AnalysisProcessingService( + analysisStateService, aiServerClient, reportGenerationService, + new AnalysisResultEnricher(), new ObjectMapper() + ); + + service.process(1L); + + verify(aiServerClient, never()).requestAnalysis(org.mockito.ArgumentMatchers.any()); + } + + @Test + void recoverStaleProcessing_doesNotCallAiWhenWorkIsNoLongerStale() { + LocalDateTime cutoff = LocalDateTime.now().minusMinutes(2); + given(analysisStateService.restartStaleProcessing(1L, cutoff)).willReturn(Optional.empty()); + AnalysisProcessingService service = new AnalysisProcessingService( + analysisStateService, aiServerClient, reportGenerationService, + new AnalysisResultEnricher(), new ObjectMapper() + ); + + service.recoverStaleProcessing(1L, cutoff); + + verify(aiServerClient, never()).requestAnalysis(org.mockito.ArgumentMatchers.any()); + } + + @Test + void process_doesNotGenerateReportWhenAiResultIsInvalid() throws Exception { + ObjectMapper objectMapper = new ObjectMapper(); + JsonNode invalidResult = objectMapper.readTree("{}"); + LocalDateTime processingStartedAt = LocalDateTime.now(); + given(analysisStateService.startProcessing(1L)) + .willReturn(Optional.of(new AnalysisProcessingClaim( + "{\"meta\":null,\"chords\":[],\"notes\":[]}", + processingStartedAt + ))); + given(aiServerClient.requestAnalysis(org.mockito.ArgumentMatchers.any(AiAnalysisRequest.class))) + .willReturn(invalidResult); + doThrow(new GeneralException(AnalysisErrorStatus.INVALID_RAW_RESULT)) + .when(analysisStateService).validateResult(invalidResult); + AnalysisProcessingService service = new AnalysisProcessingService( + analysisStateService, aiServerClient, reportGenerationService, + new AnalysisResultEnricher(), objectMapper + ); + + service.process(1L); + + verify(reportGenerationService, never()).generate(org.mockito.ArgumentMatchers.any()); + verify(analysisStateService).fail( + org.mockito.ArgumentMatchers.eq(1L), + org.mockito.ArgumentMatchers.eq(processingStartedAt), + org.mockito.ArgumentMatchers.any() + ); + } + + @Test + void process_marksMalformedRequestAsFailedWithoutCallingAi() { + LocalDateTime processingStartedAt = LocalDateTime.now(); + given(analysisStateService.startProcessing(1L)) + .willReturn(Optional.of(new AnalysisProcessingClaim("{invalid", processingStartedAt))); + AnalysisProcessingService service = new AnalysisProcessingService( + analysisStateService, aiServerClient, reportGenerationService, + new AnalysisResultEnricher(), new ObjectMapper() + ); + + service.process(1L); + + verify(aiServerClient, never()).requestAnalysis(org.mockito.ArgumentMatchers.any()); + verify(analysisStateService).fail( + org.mockito.ArgumentMatchers.eq(1L), + org.mockito.ArgumentMatchers.eq(processingStartedAt), + org.mockito.ArgumentMatchers.any() + ); + } + + @Test + void process_passesGeneratedReportToFencedCompletion() throws Exception { + ObjectMapper objectMapper = new ObjectMapper(); + LocalDateTime processingStartedAt = LocalDateTime.now(); + JsonNode validResult = objectMapper.readTree(""" + { + "scores": { + "final_score": 80, + "domains": { + "스케일": 81, + "텐션": 82, + "진행": 83, + "코드 연결": 84 + } + } + } + """); + JsonNode enrichedResult = new AnalysisResultEnricher().enrich(validResult); + GeneratedAnalysisReport generatedReport = new GeneratedAnalysisReport( + ReportGenerationType.RULE_BASED, + "리포트", + "gemini-3-flash-preview", + "analysis-report-v1", + new LlmCallMetadata( + LlmCallStatus.FAILED, + "gemini-3-flash-preview", + "analysis-report-v1", + objectMapper.createObjectNode(), + null, + null, + null, + new BigDecimal("0.30"), + 100, + false, + "input-hash", + "LLM 호출 실패" + ) + ); + given(analysisStateService.startProcessing(1L)) + .willReturn(Optional.of(new AnalysisProcessingClaim( + "{\"meta\":null,\"chords\":[],\"notes\":[]}", + processingStartedAt + ))); + given(aiServerClient.requestAnalysis(org.mockito.ArgumentMatchers.any(AiAnalysisRequest.class))) + .willReturn(validResult); + given(reportGenerationService.generate(enrichedResult)).willReturn(generatedReport); + given(analysisStateService.complete( + 1L, + processingStartedAt, + enrichedResult, + enrichedResult.toString(), + generatedReport + )).willReturn(true); + AnalysisProcessingService service = new AnalysisProcessingService( + analysisStateService, aiServerClient, reportGenerationService, + new AnalysisResultEnricher(), objectMapper + ); + + service.process(1L); + + verify(analysisStateService).validateResult(validResult); + verify(reportGenerationService).generate(enrichedResult); + verify(analysisStateService).complete( + 1L, + processingStartedAt, + enrichedResult, + enrichedResult.toString(), + generatedReport + ); + verify(analysisStateService, never()).fail( + org.mockito.ArgumentMatchers.any(), + org.mockito.ArgumentMatchers.any(), + org.mockito.ArgumentMatchers.any() + ); + } +} diff --git a/src/test/java/com/mr/domain/analysis/service/AnalysisRecoverySchedulerTest.java b/src/test/java/com/mr/domain/analysis/service/AnalysisRecoverySchedulerTest.java new file mode 100644 index 00000000..39ecd899 --- /dev/null +++ b/src/test/java/com/mr/domain/analysis/service/AnalysisRecoverySchedulerTest.java @@ -0,0 +1,88 @@ +package com.mr.domain.analysis.service; + +import static org.mockito.ArgumentMatchers.any; +import static org.mockito.ArgumentMatchers.eq; +import static org.mockito.BDDMockito.given; +import static org.mockito.Mockito.verify; + +import com.mr.domain.analysis.entity.enums.AnalysisStatus; +import com.mr.domain.analysis.repository.AnalysisRepository; +import java.time.Clock; +import java.time.Duration; +import java.time.Instant; +import java.time.LocalDateTime; +import java.time.ZoneId; +import java.util.List; + +import com.mr.domain.analysis.scheduler.AnalysisRecoveryScheduler; +import org.junit.jupiter.api.Test; +import org.junit.jupiter.api.extension.ExtendWith; +import org.mockito.Mock; +import org.mockito.junit.jupiter.MockitoExtension; +import org.springframework.core.task.SyncTaskExecutor; +import org.springframework.data.domain.Pageable; + +@ExtendWith(MockitoExtension.class) +class AnalysisRecoverySchedulerTest { + + private static final Clock FIXED_CLOCK = Clock.fixed( + Instant.parse("2026-07-31T03:00:00Z"), + ZoneId.of("Asia/Seoul") + ); + + @Mock + private AnalysisRepository analysisRepository; + + @Mock + private AnalysisProcessingService analysisProcessingService; + + @Test + void recoverPendingAnalyses_submitsPersistedPendingWork() { + given(analysisRepository.findIdsByStatusAndCreatedAtBefore( + eq(AnalysisStatus.PENDING), any(), any(Pageable.class) + )).willReturn(List.of(11L, 12L)); + given(analysisRepository.findIdsByStatusAndProcessingStartedAtBefore( + eq(AnalysisStatus.PROCESSING), any(), any(Pageable.class) + )).willReturn(List.of()); + AnalysisRecoveryScheduler scheduler = new AnalysisRecoveryScheduler( + analysisRepository, + analysisProcessingService, + new SyncTaskExecutor(), + FIXED_CLOCK, + Duration.ofMinutes(1), + Duration.ofMinutes(2), + 20 + ); + + scheduler.recoverPendingAnalyses(); + + verify(analysisProcessingService).process(11L); + verify(analysisProcessingService).process(12L); + } + + @Test + void recoverPendingAnalyses_resubmitsOnlyPersistedStaleProcessingWork() { + given(analysisRepository.findIdsByStatusAndCreatedAtBefore( + eq(AnalysisStatus.PENDING), any(), any(Pageable.class) + )).willReturn(List.of()); + given(analysisRepository.findIdsByStatusAndProcessingStartedAtBefore( + eq(AnalysisStatus.PROCESSING), any(), any(Pageable.class) + )).willReturn(List.of(21L)); + AnalysisRecoveryScheduler scheduler = new AnalysisRecoveryScheduler( + analysisRepository, + analysisProcessingService, + new SyncTaskExecutor(), + FIXED_CLOCK, + Duration.ofMinutes(1), + Duration.ofMinutes(2), + 20 + ); + + scheduler.recoverPendingAnalyses(); + + verify(analysisProcessingService).recoverStaleProcessing( + 21L, + LocalDateTime.of(2026, 7, 31, 11, 58) + ); + } +} diff --git a/src/test/java/com/mr/domain/analysis/service/AnalysisRequestFactoryTest.java b/src/test/java/com/mr/domain/analysis/service/AnalysisRequestFactoryTest.java new file mode 100644 index 00000000..7a615c39 --- /dev/null +++ b/src/test/java/com/mr/domain/analysis/service/AnalysisRequestFactoryTest.java @@ -0,0 +1,79 @@ +package com.mr.domain.analysis.service; + +import static org.assertj.core.api.Assertions.assertThat; +import static org.assertj.core.api.Assertions.assertThatThrownBy; +import static org.mockito.BDDMockito.given; +import static org.mockito.Mockito.mock; + +import com.mr.domain.analysis.exception.AnalysisErrorStatus; +import com.mr.domain.analysis.factory.AnalysisRequestFactory; +import com.mr.domain.backingTrack.entity.BackingTrack; +import com.mr.domain.backingTrack.entity.enums.Level; +import com.mr.domain.backingTrack.entity.enums.ScaleType; +import com.mr.domain.playing.entity.MidiEventData; +import com.mr.domain.playing.entity.Playing; +import com.mr.domain.playing.entity.enums.MidiType; +import com.mr.global.apipayload.exception.GeneralException; +import com.mr.global.client.ai.AiAnalysisRequest; +import java.util.List; +import org.junit.jupiter.api.BeforeEach; +import org.junit.jupiter.api.Test; + +class AnalysisRequestFactoryTest { + + private AnalysisRequestFactory factory; + private Playing playing; + private BackingTrack track; + + @BeforeEach + void setUp() { + factory = new AnalysisRequestFactory(); + playing = mock(Playing.class); + track = mock(BackingTrack.class); + given(playing.getBackingTrack()).willReturn(track); + given(playing.getBpm()).willReturn(120); + given(track.getTimeSignature()).willReturn("4/4"); + given(track.getPlaytimeSec()).willReturn(120); + given(track.getChordProgressions()).willReturn(List.of()); + given(track.getKeySignature()).willReturn("C"); + given(track.getScaleType()).willReturn(ScaleType.MAJOR); + given(track.getGenre()).willReturn("JAZZ"); + given(track.getLevel()).willReturn(Level.BASIC); + } + + @Test + void create_slicesAndRebasesMidiByBar() { + given(playing.getMidiData()).willReturn(List.of( + MidiEventData.of(0, MidiType.NOTE_ON, 60, 100, 0L), + MidiEventData.of(1, MidiType.NOTE_ON, 62, 100, 2_000L), + MidiEventData.of(2, MidiType.NOTE_OFF, 62, 0, 3_999L), + MidiEventData.of(3, MidiType.NOTE_ON, 64, 100, 4_000L) + )); + + AiAnalysisRequest request = factory.create(playing, 2, 2); + + assertThat(request.notes()).hasSize(2); + assertThat(request.notes().get(0).type()).isEqualTo(AiAnalysisRequest.NoteType.NOTE_ON); + assertThat(request.notes().get(1).type()).isEqualTo(AiAnalysisRequest.NoteType.NOTE_OFF); + assertThat(request.notes().get(0).timestampMs()).isEqualTo(0D); + assertThat(request.notes().get(1).timestampMs()).isEqualTo(1_999D); + } + + @Test + void create_rejectsMoreThanThirtyTwoBars() { + assertThatThrownBy(() -> factory.create(playing, 1, 33)) + .isInstanceOf(GeneralException.class) + .hasFieldOrPropertyWithValue("code", AnalysisErrorStatus.INVALID_BAR_RANGE); + } + + @Test + void create_rejectsValidBarRangeWithoutNotes() { + given(playing.getMidiData()).willReturn(List.of( + MidiEventData.of(0, MidiType.NOTE_ON, 60, 100, 0L) + )); + + assertThatThrownBy(() -> factory.create(playing, 2, 2)) + .isInstanceOf(GeneralException.class) + .hasFieldOrPropertyWithValue("code", AnalysisErrorStatus.EMPTY_NOTE_RANGE); + } +} diff --git a/src/test/java/com/mr/domain/analysis/service/AnalysisServiceTest.java b/src/test/java/com/mr/domain/analysis/service/AnalysisServiceTest.java index 4ee5b32c..e58c7bf6 100644 --- a/src/test/java/com/mr/domain/analysis/service/AnalysisServiceTest.java +++ b/src/test/java/com/mr/domain/analysis/service/AnalysisServiceTest.java @@ -12,26 +12,41 @@ import static org.assertj.core.api.Assertions.assertThatThrownBy; import com.fasterxml.jackson.databind.ObjectMapper; +import com.mr.domain.analysis.dto.req.AnalysisCreateRequestDTO; +import com.mr.domain.analysis.dto.res.AnalysisCreateResponseDTO; import com.mr.domain.analysis.dto.res.AnalysisResultResponseDTO; import com.mr.domain.analysis.entity.Analysis; import com.mr.domain.analysis.entity.AnalysisReport; import com.mr.domain.analysis.entity.enums.AnalysisGrade; import com.mr.domain.analysis.entity.enums.LlmStatus; +import com.mr.domain.analysis.entity.enums.AnalysisStatus; +import com.mr.domain.analysis.event.AnalysisRequestedEvent; import com.mr.domain.analysis.exception.AnalysisErrorStatus; +import com.mr.domain.analysis.factory.AnalysisRequestFactory; import com.mr.domain.analysis.repository.AnalysisReportRepository; import com.mr.domain.analysis.repository.AnalysisRepository; +import com.mr.domain.backingTrack.entity.BackingTrack; +import com.mr.domain.backingTrack.entity.enums.ScaleType; import com.mr.domain.playing.entity.Playing; +import com.mr.domain.playing.entity.enums.PlayingStatus; +import com.mr.domain.playing.repository.PlayingRepository; import com.mr.domain.user.entity.User; import com.mr.global.apipayload.exception.GeneralException; +import com.mr.global.client.ai.AiAnalysisRequest; import java.math.BigDecimal; +import java.time.LocalDateTime; +import java.util.List; +import java.util.Locale; import java.util.Optional; import org.junit.jupiter.api.BeforeEach; import org.junit.jupiter.api.DisplayName; import org.junit.jupiter.api.Test; import org.junit.jupiter.api.extension.ExtendWith; import org.mockito.ArgumentCaptor; +import org.mockito.stubbing.Answer; import org.mockito.Mock; import org.mockito.junit.jupiter.MockitoExtension; +import org.springframework.context.ApplicationEventPublisher; @ExtendWith(MockitoExtension.class) class AnalysisServiceTest { @@ -42,23 +57,45 @@ class AnalysisServiceTest { @Mock private AnalysisReportRepository analysisReportRepository; + @Mock + private PlayingRepository playingRepository; + + @Mock + private AnalysisRequestFactory analysisRequestFactory; + + @Mock + private ApplicationEventPublisher eventPublisher; + private AnalysisService analysisService; @BeforeEach void setUp() { - analysisService = new AnalysisService(analysisRepository, analysisReportRepository, new ObjectMapper()); + analysisService = new AnalysisService(analysisRepository, analysisReportRepository, playingRepository, + analysisRequestFactory, eventPublisher, new ObjectMapper()); } private Analysis completedAnalysis(Long userId) { + return completedAnalysis(userId, ScaleType.MAJOR); + } + + private Analysis completedAnalysis(Long userId, ScaleType scaleType) { User user = mock(User.class); given(user.getUserId()).willReturn(userId); Playing playing = mock(Playing.class); + BackingTrack backingTrack = mock(BackingTrack.class); lenient().when(playing.getId()).thenReturn(1L); lenient().when(playing.getUser()).thenReturn(user); + lenient().when(playing.getBackingTrack()).thenReturn(backingTrack); + lenient().when(playing.getBpm()).thenReturn(120); + lenient().when(backingTrack.getTitle()).thenReturn("테스트 트랙"); + lenient().when(backingTrack.getGenre()).thenReturn("jazz"); + lenient().when(backingTrack.getKeySignature()).thenReturn("C"); + lenient().when(backingTrack.getScaleType()).thenReturn(scaleType); Analysis analysis = Analysis.createPending(user, playing, 1, 8, "{}"); - analysis.startProcessing(); + LocalDateTime now = LocalDateTime.of(2026, 7, 31, 12, 0); + analysis.startProcessing(now); analysis.complete( 85, AnalysisGrade.GOOD, @@ -67,11 +104,31 @@ private Analysis completedAnalysis(Long userId) { new BigDecimal("75.50"), new BigDecimal("90.00"), new BigDecimal("70.00"), - null + null, + now ); return analysis; } + @Test + @DisplayName("getAnalysisResult - JVM 로케일과 무관하게 조성명을 변환한다") + void getAnalysisResult_turkishLocale_formatsKeyConsistently() { + Locale previousLocale = Locale.getDefault(); + try { + Locale.setDefault(Locale.forLanguageTag("tr-TR")); + Analysis analysis = completedAnalysis(1L, ScaleType.MINOR); + given(analysisRepository.findById(1L)).willReturn(Optional.of(analysis)); + given(analysisReportRepository.findFirstByAnalysisIdAndLlmStatusOrderByCreatedAtDesc(anyLong(), any())) + .willReturn(Optional.empty()); + + AnalysisResultResponseDTO response = analysisService.getAnalysisResult(1L, 1L); + + assertThat(response.key()).isEqualTo("C Minor"); + } finally { + Locale.setDefault(previousLocale); + } + } + @Test @DisplayName("getAnalysisResult - SUCCESS 리포트가 있으면 report를 채워서 반환한다") void getAnalysisResult_success_returnsLatestSuccessReport() { @@ -94,6 +151,8 @@ void getAnalysisResult_success_returnsLatestSuccessReport() { assertThat(response.report()).isNotNull(); assertThat(response.report().content()).isEqualTo("리포트 본문"); assertThat(response.report().llmStatus()).isEqualTo(LlmStatus.SUCCESS); + assertThat(response.title()).isEqualTo("테스트 트랙"); + assertThat(response.key()).isEqualTo("C Major"); } @Test @@ -127,4 +186,38 @@ void getAnalysisResult_otherUsersAnalysis_throwsAccessDenied() { .isInstanceOf(GeneralException.class) .hasFieldOrPropertyWithValue("code", AnalysisErrorStatus.ANALYSIS_ACCESS_DENIED); } + @Test + @DisplayName("createAnalysis - saves a pending analysis and publishes an event") + void createAnalysis_success() { + User user = mock(User.class); + Playing playing = mock(Playing.class); + given(user.getUserId()).willReturn(1L); + given(playing.getId()).willReturn(31L); + given(playing.getUser()).willReturn(user); + given(playing.getStatus()).willReturn(PlayingStatus.COMPLETED); + given(playingRepository.findByIdWithBackingTrackForUpdate(31L)).willReturn(Optional.of(playing)); + given(analysisRepository.existsByPlayingIdAndStatusIn(eq(31L), any())).willReturn(false); + given(analysisRequestFactory.create(playing, 1, 8)) + .willReturn(new AiAnalysisRequest(null, List.of(), List.of())); + given(analysisRepository.save(any(Analysis.class))) + .willAnswer((Answer) invocation -> invocation.getArgument(0)); + + AnalysisCreateResponseDTO response = analysisService.createAnalysis( + 1L, new AnalysisCreateRequestDTO(31L, 1, 8) + ); + + verify(analysisRepository).existsByPlayingIdAndStatusIn( + 31L, + List.of( + AnalysisStatus.PENDING, + AnalysisStatus.PROCESSING, + AnalysisStatus.COMPLETED + ) + ); + assertThat(response.playingId()).isEqualTo(31L); + assertThat(response.status()).isEqualTo(AnalysisStatus.PENDING); + verify(eventPublisher).publishEvent(any(AnalysisRequestedEvent.class)); + } + + } diff --git a/src/test/java/com/mr/domain/analysis/service/AnalysisStateServiceTest.java b/src/test/java/com/mr/domain/analysis/service/AnalysisStateServiceTest.java new file mode 100644 index 00000000..42087627 --- /dev/null +++ b/src/test/java/com/mr/domain/analysis/service/AnalysisStateServiceTest.java @@ -0,0 +1,352 @@ +package com.mr.domain.analysis.service; + +import static org.assertj.core.api.Assertions.assertThatThrownBy; +import static org.mockito.ArgumentMatchers.any; +import static org.mockito.BDDMockito.given; +import static org.mockito.Mockito.mock; +import static org.mockito.Mockito.never; +import static org.mockito.Mockito.verify; + +import com.fasterxml.jackson.databind.JsonNode; +import com.fasterxml.jackson.databind.ObjectMapper; +import com.mr.domain.analysis.entity.Analysis; +import com.mr.domain.analysis.entity.AnalysisReport; +import com.mr.domain.analysis.entity.enums.AnalysisStatus; +import com.mr.domain.analysis.entity.enums.ReportGenerationType; +import com.mr.domain.analysis.exception.AnalysisErrorStatus; +import com.mr.domain.analysis.model.GeneratedAnalysisReport; +import com.mr.domain.analysis.model.LlmCallMetadata; +import com.mr.domain.analysis.repository.AnalysisReportRepository; +import com.mr.domain.analysis.repository.AnalysisRepository; +import com.mr.domain.mentor.entity.enums.LlmCallStatus; +import com.mr.domain.mentor.entity.LlmCallLog; +import com.mr.domain.mentor.repository.LlmCallLogRepository; +import com.mr.domain.user.entity.User; +import com.mr.global.apipayload.exception.GeneralException; +import com.mr.global.event.AnalysisCompletedEvent; +import java.math.BigDecimal; +import java.time.Clock; +import java.time.Instant; +import java.time.LocalDateTime; +import java.time.ZoneId; +import java.util.Optional; +import org.junit.jupiter.api.BeforeEach; +import org.junit.jupiter.api.Test; +import org.junit.jupiter.api.extension.ExtendWith; +import org.mockito.ArgumentCaptor; +import org.mockito.Mock; +import org.mockito.junit.jupiter.MockitoExtension; +import org.springframework.context.ApplicationEventPublisher; + +@ExtendWith(MockitoExtension.class) +class AnalysisStateServiceTest { + + private static final LocalDateTime PROCESSING_STARTED_AT = LocalDateTime.of(2026, 7, 31, 12, 0); + private static final Clock FIXED_CLOCK = Clock.fixed( + Instant.parse("2026-07-31T03:00:00Z"), + ZoneId.of("Asia/Seoul") + ); + + @Mock + private AnalysisRepository analysisRepository; + + @Mock + private AnalysisReportRepository analysisReportRepository; + + @Mock + private LlmCallLogRepository llmCallLogRepository; + + @Mock + private ApplicationEventPublisher eventPublisher; + + private AnalysisStateService service; + private Analysis analysis; + private User user; + private ObjectMapper objectMapper; + + @BeforeEach + void setUp() { + service = new AnalysisStateService( + analysisRepository, + analysisReportRepository, + llmCallLogRepository, + eventPublisher, + FIXED_CLOCK + ); + analysis = mock(Analysis.class); + user = mock(User.class); + org.mockito.Mockito.lenient().when(user.getUserId()).thenReturn(1L); + org.mockito.Mockito.lenient().when(analysis.getUser()).thenReturn(user); + objectMapper = new ObjectMapper(); + org.mockito.Mockito.lenient() + .when(analysisRepository.findByIdForUpdate(1L)) + .thenReturn(Optional.of(analysis)); + org.mockito.Mockito.lenient() + .when(analysis.isCurrentProcessing(PROCESSING_STARTED_AT)) + .thenReturn(true); + } + + @Test + void startProcessing_marksBlankRequestAsFailed() { + given(analysis.getStatus()).willReturn(AnalysisStatus.PENDING); + given(analysis.getAnalysisRequestJson()).willReturn(" "); + + Optional claim = service.startProcessing(1L); + + org.assertj.core.api.Assertions.assertThat(claim).isEmpty(); + verify(analysis).fail(AnalysisErrorStatus.INVALID_ANALYSIS_REQUEST.getMessage(), PROCESSING_STARTED_AT); + verify(analysis, never()).startProcessing(any()); + } + + @Test + void restartStaleProcessing_marksNullRequestAsFailed() { + LocalDateTime cutoff = LocalDateTime.now(); + given(analysis.getStatus()).willReturn(AnalysisStatus.PROCESSING); + given(analysis.getProcessingStartedAt()).willReturn(cutoff.minusMinutes(1)); + given(analysis.getAnalysisRequestJson()).willReturn(null); + + Optional claim = service.restartStaleProcessing(1L, cutoff); + + org.assertj.core.api.Assertions.assertThat(claim).isEmpty(); + verify(analysis).fail(AnalysisErrorStatus.INVALID_ANALYSIS_REQUEST.getMessage(), PROCESSING_STARTED_AT); + verify(analysis, never()).restartProcessing(any()); + } + + @Test + void complete_acceptsCompleteScoresWithinRange() throws Exception { + JsonNode result = result(""" + { + "scores": { + "final_score": 80.5, + "grade": "GOOD", + "domains": { + "스케일": 81.0, + "텐션": 82.0, + "진행": 83.0, + "코드 연결": 84.0 + } + }, + "summary": "요약" + } + """); + + service.complete(1L, PROCESSING_STARTED_AT, result, result.toString(), generatedReport()); + + verify(analysis).complete( + 81, + com.mr.domain.analysis.entity.enums.AnalysisGrade.GOOD, + "요약", + new BigDecimal("81.0"), + new BigDecimal("82.0"), + new BigDecimal("83.0"), + new BigDecimal("84.0"), + result.toString(), + PROCESSING_STARTED_AT + ); + verify(analysisReportRepository).save(any(AnalysisReport.class)); + verify(llmCallLogRepository).save(any(LlmCallLog.class)); + verify(eventPublisher).publishEvent( + org.mockito.ArgumentMatchers.argThat(event -> + event instanceof AnalysisCompletedEvent completedEvent + && completedEvent.getUserId().equals(1L) + ) + ); + } + + @Test + void complete_persistsLlmReportAndSuccessfulCallLog() throws Exception { + JsonNode result = validResult(); + + service.complete( + 1L, + PROCESSING_STARTED_AT, + result, + result.toString(), + generatedReport(ReportGenerationType.LLM, LlmCallStatus.SUCCESS) + ); + + ArgumentCaptor reportCaptor = ArgumentCaptor.forClass(AnalysisReport.class); + ArgumentCaptor logCaptor = ArgumentCaptor.forClass(LlmCallLog.class); + verify(analysisReportRepository).save(reportCaptor.capture()); + verify(llmCallLogRepository).save(logCaptor.capture()); + org.assertj.core.api.Assertions.assertThat(reportCaptor.getValue().getGenerationType()) + .isEqualTo(ReportGenerationType.LLM); + org.assertj.core.api.Assertions.assertThat(logCaptor.getValue().getStatus()) + .isEqualTo(LlmCallStatus.SUCCESS); + org.assertj.core.api.Assertions.assertThat(logCaptor.getValue().getTotalTokens()).isEqualTo(30); + } + + @Test + void complete_persistsTimeoutCallLogForFallbackReport() throws Exception { + JsonNode result = validResult(); + + service.complete( + 1L, + PROCESSING_STARTED_AT, + result, + result.toString(), + generatedReport(ReportGenerationType.RULE_BASED, LlmCallStatus.TIMEOUT) + ); + + ArgumentCaptor logCaptor = ArgumentCaptor.forClass(LlmCallLog.class); + verify(llmCallLogRepository).save(logCaptor.capture()); + org.assertj.core.api.Assertions.assertThat(logCaptor.getValue().getStatus()) + .isEqualTo(LlmCallStatus.TIMEOUT); + org.assertj.core.api.Assertions.assertThat(logCaptor.getValue().getErrorMessage()) + .isEqualTo("failed"); + } + + @Test + void complete_ignoresStaleProcessingAttempt() throws Exception { + JsonNode result = result(""" + { + "scores": { + "final_score": 80, + "domains": { + "스케일": 81, + "텐션": 82, + "진행": 83, + "코드 연결": 84 + } + } + } + """); + given(analysis.isCurrentProcessing(PROCESSING_STARTED_AT)).willReturn(false); + + boolean completed = service.complete( + 1L, PROCESSING_STARTED_AT, result, result.toString(), generatedReport() + ); + + org.assertj.core.api.Assertions.assertThat(completed).isFalse(); + verify(analysis, never()).complete(any(), any(), any(), any(), any(), any(), any(), any(), any()); + verify(analysisReportRepository, never()).save(any()); + verify(llmCallLogRepository, never()).save(any()); + verify(eventPublisher, never()).publishEvent(any()); + } + + @Test + void fail_ignoresStaleProcessingAttempt() { + given(analysis.isCurrentProcessing(PROCESSING_STARTED_AT)).willReturn(false); + + boolean failed = service.fail(1L, PROCESSING_STARTED_AT, "timeout"); + + org.assertj.core.api.Assertions.assertThat(failed).isFalse(); + verify(analysis, never()).fail(any(), any()); + } + + @Test + void complete_rejectsMissingDomainScore() throws Exception { + JsonNode result = result(""" + { + "scores": { + "final_score": 80, + "domains": { + "스케일": 81, + "텐션": 82, + "진행": 83 + } + } + } + """); + + assertInvalidRawResult(result); + } + + @Test + void complete_rejectsScoreOutsideZeroToOneHundred() throws Exception { + JsonNode result = result(""" + { + "scores": { + "final_score": 101, + "domains": { + "스케일": 81, + "텐션": 82, + "진행": 83, + "코드 연결": 84 + } + } + } + """); + + assertInvalidRawResult(result); + } + + @Test + void complete_rejectsNonNumericScore() throws Exception { + JsonNode result = result(""" + { + "scores": { + "final_score": 80, + "domains": { + "스케일": "81", + "텐션": 82, + "진행": 83, + "코드 연결": 84 + } + } + } + """); + + assertInvalidRawResult(result); + } + + private JsonNode result(String json) throws Exception { + return objectMapper.readTree(json); + } + + private JsonNode validResult() throws Exception { + return result(""" + { + "scores": { + "final_score": 80, + "grade": "GOOD", + "domains": { + "스케일": 81, + "텐션": 82, + "진행": 83, + "코드 연결": 84 + } + } + } + """); + } + + private void assertInvalidRawResult(JsonNode result) { + assertThatThrownBy(() -> service.complete( + 1L, PROCESSING_STARTED_AT, result, result.toString(), generatedReport() + )) + .isInstanceOf(GeneralException.class) + .hasFieldOrPropertyWithValue("code", AnalysisErrorStatus.INVALID_RAW_RESULT); + verify(analysis, never()).complete(any(), any(), any(), any(), any(), any(), any(), any(), any()); + } + + private GeneratedAnalysisReport generatedReport() { + return generatedReport(ReportGenerationType.RULE_BASED, LlmCallStatus.FAILED); + } + + private GeneratedAnalysisReport generatedReport( + ReportGenerationType generationType, + LlmCallStatus status + ) { + return new GeneratedAnalysisReport( + generationType, + "리포트", + "gemini-3-flash-preview", + "analysis-report-v1", + new LlmCallMetadata( + status, + "gemini-3-flash-preview", + "analysis-report-v1", + objectMapper.createObjectNode(), + 10, + 20, + 30, + new BigDecimal("0.30"), + 100, + false, + "hash", + "failed" + ) + ); + } +} diff --git a/src/test/java/com/mr/domain/analysis/service/ReportGenerationServiceTest.java b/src/test/java/com/mr/domain/analysis/service/ReportGenerationServiceTest.java new file mode 100644 index 00000000..1e560080 --- /dev/null +++ b/src/test/java/com/mr/domain/analysis/service/ReportGenerationServiceTest.java @@ -0,0 +1,197 @@ +package com.mr.domain.analysis.service; + +import static org.assertj.core.api.Assertions.assertThat; +import static org.mockito.BDDMockito.given; + +import com.fasterxml.jackson.databind.JsonNode; +import com.fasterxml.jackson.databind.ObjectMapper; +import com.mr.domain.analysis.entity.enums.ReportGenerationType; +import com.mr.domain.analysis.generator.RuleBasedReportGenerator; +import com.mr.domain.analysis.model.GeneratedAnalysisReport; +import com.mr.global.client.gemini.GeminiClient; +import com.mr.global.client.gemini.GeminiGenerationResult; +import com.mr.global.config.GeminiProperties; +import java.time.Duration; +import org.junit.jupiter.api.BeforeEach; +import org.junit.jupiter.api.Test; +import org.junit.jupiter.api.extension.ExtendWith; +import org.mockito.Mock; +import org.mockito.junit.jupiter.MockitoExtension; + +@ExtendWith(MockitoExtension.class) +class ReportGenerationServiceTest { + + @Mock + private GeminiClient geminiClient; + + private ReportGenerationService service; + private JsonNode result; + + @BeforeEach + void setUp() throws Exception { + GeminiProperties properties = new GeminiProperties( + "https://generativelanguage.googleapis.com", + "key", + "gemini-3-flash-preview", + Duration.ofSeconds(5), + Duration.ofSeconds(60) + ); + service = new ReportGenerationService( + geminiClient, + properties, + new RuleBasedReportGenerator(), + new ObjectMapper() + ); + result = new ObjectMapper().readTree(""" + { + "meta": {"key":"C major","genre":"jazz","time_signature":[4,4],"bpm":120}, + "scores": { + "final_score":80, + "domains":{"스케일":90,"텐션":70,"진행":85,"코드 연결":75} + } + } + """); + } + + @Test + void generate_returnsLlmReportWhenGeminiSucceeds() { + given(geminiClient.generateReport( + org.mockito.ArgumentMatchers.anyString(), + org.mockito.ArgumentMatchers.anyString() + )).willReturn(new GeminiGenerationResult( + validMarkdownReport(), 100, 50, 150, false + )); + + GeneratedAnalysisReport report = service.generate(result); + + assertThat(report.generationType()).isEqualTo(ReportGenerationType.LLM); + assertThat(report.content()).isEqualTo(validMarkdownReport()); + assertThat(report.modelName()).isEqualTo("gemini-3-flash-preview"); + assertThat(report.llmCall().promptTokens()).isEqualTo(100); + assertThat(report.llmCall().totalTokens()).isEqualTo(150); + org.mockito.Mockito.verify(geminiClient).generateReport( + org.mockito.ArgumentMatchers.argThat(prompt -> + prompt.contains("700자 이상 1,500자 이하") + && prompt.contains("문제점·근거·실행 가능한 연습 방법") + ), + org.mockito.ArgumentMatchers.anyString() + ); + } + + @Test + void generate_fallsBackWhenGeminiResponseIsMissingRequiredHeading() { + given(geminiClient.generateReport( + org.mockito.ArgumentMatchers.anyString(), + org.mockito.ArgumentMatchers.anyString() + )).willReturn(new GeminiGenerationResult( + """ + # 연주 분석 리포트 + ## 총평 + ## 잘한 점 + ## 진행 맥락 + ## 점수 요약 + """, + 100, 50, 150, false + )); + + GeneratedAnalysisReport report = service.generate(result); + + assertThat(report.generationType()).isEqualTo(ReportGenerationType.RULE_BASED); + assertThat(report.content()).contains("# 연주 분석 리포트", "## 개선 제안", "80 / 100"); + assertThat(report.llmCall().status()) + .isEqualTo(com.mr.domain.mentor.entity.enums.LlmCallStatus.FAILED); + } + + @Test + void generate_fallsBackWhenGeminiHeadingsAreOutOfOrder() { + given(geminiClient.generateReport( + org.mockito.ArgumentMatchers.anyString(), + org.mockito.ArgumentMatchers.anyString() + )).willReturn(new GeminiGenerationResult( + """ + # 연주 분석 리포트 + ## 잘한 점 + ## 총평 + ## 진행 맥락 + ## 개선 제안 + ## 점수 요약 + """, + 100, 50, 150, false + )); + + GeneratedAnalysisReport report = service.generate(result); + + assertThat(report.generationType()).isEqualTo(ReportGenerationType.RULE_BASED); + assertThat(report.content()).contains("# 연주 분석 리포트", "## 총평", "## 잘한 점"); + } + + @Test + void generate_fallsBackWhenGeminiReportIsStructurallyValidButTooShort() { + given(geminiClient.generateReport( + org.mockito.ArgumentMatchers.anyString(), + org.mockito.ArgumentMatchers.anyString() + )).willReturn(new GeminiGenerationResult( + """ + # 연주 분석 리포트 + ## 총평 + 짧은 총평 + ## 잘한 점 + 짧은 강점 + ## 진행 맥락 + 짧은 맥락 + ## 개선 제안 + 짧은 제안 + ## 점수 요약 + 80점 + """, + 100, 50, 150, false + )); + + GeneratedAnalysisReport report = service.generate(result); + + assertThat(report.generationType()).isEqualTo(ReportGenerationType.RULE_BASED); + assertThat(report.llmCall().status()) + .isEqualTo(com.mr.domain.mentor.entity.enums.LlmCallStatus.FAILED); + } + + @Test + void generate_fallsBackToRuleBasedReportWhenGeminiFails() { + given(geminiClient.generateReport( + org.mockito.ArgumentMatchers.anyString(), + org.mockito.ArgumentMatchers.anyString() + )).willThrow(new RuntimeException("quota exceeded")); + + GeneratedAnalysisReport report = service.generate(result); + + assertThat(report.generationType()).isEqualTo(ReportGenerationType.RULE_BASED); + assertThat(report.content()).contains("# 연주 분석 리포트", "## 개선 제안", "80 / 100"); + assertThat(report.modelName()).isNull(); + assertThat(report.llmCall().status()) + .isEqualTo(com.mr.domain.mentor.entity.enums.LlmCallStatus.FAILED); + } + + private String validMarkdownReport() { + return """ + # 연주 분석 리포트 + **조성** C major · **장르** jazz · **박자** 4/4 · **템포** 120 bpm + ## 총평 + 종합 점수는 80점으로 전반적인 코드 진행 이해도가 안정적입니다. 스케일 점수가 가장 높아 조성 안에서 음을 선택하는 능력이 잘 드러났습니다. 진행 점수 역시 안정적이어서 백킹트랙의 흐름을 크게 벗어나지 않았습니다. 다만 텐션과 코드 연결 점수는 상대적으로 낮으므로 다음 연습에서 우선 확인할 필요가 있습니다. + ## 잘한 점 + - 스케일 영역은 90점으로 네 영역 중 가장 높습니다. 입력 결과에서 확인된 조성 안의 음을 일관되게 선택한 점이 강점입니다. + - 진행 영역은 85점입니다. 코드가 바뀌는 구간에서도 전체 화성 흐름을 유지하여 연주의 맥락이 끊기지 않았습니다. + ## 진행 맥락 + - 입력에 기록된 코드 진행을 따라 연주가 이어졌으며, 진행 점수 85점이 이러한 일관성을 뒷받침합니다. + - 코드 연결은 75점으로 기본 흐름은 유지했지만, 다음 코드로 이동할 때 더 가까운 음을 선택할 여지가 있습니다. + ## 개선 제안 + - 텐션 영역은 70점으로 가장 낮습니다. 먼저 코드톤을 확인한 뒤 9음이나 13음을 한 종류씩 추가하여 색채 변화를 비교해 보세요. + - 코드 연결 영역은 75점입니다. 같은 진행을 느린 템포로 반복하면서 이전 코드의 마지막 음과 다음 코드의 첫 음 사이 간격을 줄여 보세요. + - 점수를 높이기 위해 빠르게 반복하기보다, 각 코드에서 선택한 음이 코드톤인지 텐션인지 소리로 확인하는 연습이 적합합니다. + ## 점수 요약 + - 종합 점수: 80 / 100 + - 스케일: 90 + - 텐션: 70 + - 진행: 85 + - 코드 연결: 75 + """; + } +} diff --git a/src/test/java/com/mr/domain/history/service/HistoryServiceTest.java b/src/test/java/com/mr/domain/history/service/HistoryServiceTest.java index 75bdaf7b..50fb377e 100644 --- a/src/test/java/com/mr/domain/history/service/HistoryServiceTest.java +++ b/src/test/java/com/mr/domain/history/service/HistoryServiceTest.java @@ -74,8 +74,9 @@ private Analysis completedAnalysis(Long playingId, Integer totalScore) { lenient().when(playing.getUser()).thenReturn(user); Analysis analysis = Analysis.createPending(user, playing, 1, 8, "{}"); - analysis.startProcessing(); - analysis.complete(totalScore, AnalysisGrade.GOOD, "요약", null, null, null, null, null); + LocalDateTime now = LocalDateTime.of(2026, 7, 31, 12, 0); + analysis.startProcessing(now); + analysis.complete(totalScore, AnalysisGrade.GOOD, "요약", null, null, null, null, null, now); return analysis; } diff --git a/src/test/java/com/mr/domain/statistics/service/StatisticsAggregationServiceTest.java b/src/test/java/com/mr/domain/statistics/service/StatisticsAggregationServiceTest.java index 0583328c..212a1c20 100644 --- a/src/test/java/com/mr/domain/statistics/service/StatisticsAggregationServiceTest.java +++ b/src/test/java/com/mr/domain/statistics/service/StatisticsAggregationServiceTest.java @@ -373,7 +373,6 @@ void onAnalysisCompleted_skillScoreAllNull_skipsThatSkillType() { .willReturn(emptyAnalysisTotals); given(userStatisticsRepository.findByUser_UserId(userId)) .willReturn(Optional.of(UserStatistics.createForUser(mock(User.class)))); - // scaleScore만 null(AI가 해당 지표를 안 준 경우), 나머지는 정상 값 List weeklyAnalyses = List.of(mockAnalysis(null, null, new BigDecimal("70.0"), new BigDecimal("80.0"), new BigDecimal("90.0"))); given(analysisRepository.findByUserAndStatusSince(userId, AnalysisStatus.COMPLETED, weekStart.atStartOfDay())) diff --git a/src/test/java/com/mr/domain/statistics/service/StatisticsEventListenerRetryTest.java b/src/test/java/com/mr/domain/statistics/service/StatisticsEventListenerRetryTest.java index 2c07552e..5a013cf2 100644 --- a/src/test/java/com/mr/domain/statistics/service/StatisticsEventListenerRetryTest.java +++ b/src/test/java/com/mr/domain/statistics/service/StatisticsEventListenerRetryTest.java @@ -18,8 +18,7 @@ import org.springframework.test.context.junit.jupiter.SpringExtension; import org.springframework.beans.factory.annotation.Autowired; -// @Retryable/@Recover는 AOP 프록시를 거쳐야만 동작하므로, Mockito만으로 new해서 직접 호출하면 -// 재시도 자체가 발생하지 않아 검증할 수 없다. 최소한의 스프링 컨텍스트로 프록시를 실제로 태운다. +// 재시도 검증용 AOP 프록시 컨텍스트 @ExtendWith(SpringExtension.class) @ContextConfiguration(classes = { StatisticsEventListener.class, diff --git a/src/test/java/com/mr/domain/statistics/service/StatisticsServiceTest.java b/src/test/java/com/mr/domain/statistics/service/StatisticsServiceTest.java index c5c2bae7..6c3f1604 100644 --- a/src/test/java/com/mr/domain/statistics/service/StatisticsServiceTest.java +++ b/src/test/java/com/mr/domain/statistics/service/StatisticsServiceTest.java @@ -159,17 +159,17 @@ void getStatistics_weeklySummary_countsFromPlayingOnly() { List playings = List.of( mockPlaying(thisWeekStart.plusDays(1), 600), mockPlaying(thisWeekStart.plusDays(2), 600), - mockPlaying(thisWeekStart.minusDays(3), 1200) // 지난 주 + mockPlaying(thisWeekStart.minusDays(3), 1200) ); given(playingRepository.findByUserAndStatusSince(eq(1L), eq(PlayingStatus.COMPLETED), any())) .willReturn(playings); StatisticsResponseDTO response = statisticsService.getStatistics(1L); - assertThat(response.weeklySummary().practiceMinutes()).isEqualTo(20); // (600+600)초 = 20분 + assertThat(response.weeklySummary().practiceMinutes()).isEqualTo(20); assertThat(response.weeklySummary().completedSessionCount()).isEqualTo(2); - assertThat(response.weeklySummary().practiceMinutesDiff()).isEqualTo(0); // 20 - 20(1200초=20분) - assertThat(response.weeklySummary().completedSessionCountDiff()).isEqualTo(1); // 2 - 1 + assertThat(response.weeklySummary().practiceMinutesDiff()).isEqualTo(0); + assertThat(response.weeklySummary().completedSessionCountDiff()).isEqualTo(1); } @Test @@ -271,7 +271,7 @@ void getStatistics_playingVsAnalysisCount_areIndependent() { StatisticsResponseDTO response = statisticsService.getStatistics(1L); assertThat(response.weeklySummary().completedSessionCount()).isEqualTo(1); - assertThat(response.weeklySummary().accuracy()).isEqualByComparingTo(BigDecimal.valueOf(90.0)); // (80+90+100)/3 + assertThat(response.weeklySummary().accuracy()).isEqualByComparingTo(BigDecimal.valueOf(90.0)); } @Test @@ -298,7 +298,7 @@ void getStatistics_weeklyTrend_fourItemsInFixedOrder() { assertThat(items.get(3).label()).isEqualTo("이번주"); assertThat(items.get(0).averageScore()).isEqualByComparingTo(BigDecimal.valueOf(60.0)); assertThat(items.get(3).averageScore()).isEqualByComparingTo(BigDecimal.valueOf(93.0)); - assertThat(response.weeklyTrend().diffFromPreviousWeek()).isEqualTo(30); // 93 - 63 + assertThat(response.weeklyTrend().diffFromPreviousWeek()).isEqualTo(30); } @Test diff --git a/src/test/java/com/mr/global/client/ai/AiAnalysisRequestSerializationTest.java b/src/test/java/com/mr/global/client/ai/AiAnalysisRequestSerializationTest.java index a6dbe55d..af62b37e 100644 --- a/src/test/java/com/mr/global/client/ai/AiAnalysisRequestSerializationTest.java +++ b/src/test/java/com/mr/global/client/ai/AiAnalysisRequestSerializationTest.java @@ -16,24 +16,30 @@ class AiAnalysisRequestSerializationTest { void 요청_직렬화시_AI_서버_스키마와_동일한_snake_case_필드명을_사용한다() throws Exception { AiAnalysisRequest request = new AiAnalysisRequest( new AiAnalysisRequest.Meta(120.0, List.of(4, 4), - new AiAnalysisRequest.Key("C", "major"), "jazz"), + new AiAnalysisRequest.Key("C", "major"), "jazz", "basic"), List.of(new AiAnalysisRequest.Chord(1, 1.0, "Dm7")), - List.of(new AiAnalysisRequest.Note(0, 62, 0.0, 1.0, 90)) + List.of(new AiAnalysisRequest.Note(0, AiAnalysisRequest.NoteType.NOTE_ON, 62, 90, 0.0)) ); JsonNode json = objectMapper.valueToTree(request); assertThat(json.path("meta").has("time_signature")).isTrue(); assertThat(json.path("meta").has("timeSignature")).isFalse(); - assertThat(json.path("notes").get(0).has("onset_beats")).isTrue(); - assertThat(json.path("notes").get(0).has("duration_beats")).isTrue(); + assertThat(json.path("meta").has("level")).isTrue(); + assertThat(json.path("notes").get(0).has("timestamp_ms")).isTrue(); + assertThat(json.path("notes").get(0).has("timestampMs")).isFalse(); + assertThat(json.path("notes").get(0).has("onset_beats")).isFalse(); + assertThat(json.path("notes").get(0).has("duration_beats")).isFalse(); assertThat(json.path("meta").path("bpm").asDouble()).isEqualTo(120.0); assertThat(json.path("meta").path("key").path("tonic").asText()).isEqualTo("C"); assertThat(json.path("meta").path("key").path("mode").asText()).isEqualTo("major"); + assertThat(json.path("meta").path("level").asText()).isEqualTo("basic"); assertThat(json.path("chords").get(0).path("bar").asInt()).isEqualTo(1); assertThat(json.path("chords").get(0).path("symbol").asText()).isEqualTo("Dm7"); assertThat(json.path("notes").get(0).path("pitch").asInt()).isEqualTo(62); assertThat(json.path("notes").get(0).path("velocity").asInt()).isEqualTo(90); + assertThat(json.path("notes").get(0).path("type").asText()).isEqualTo("NOTE_ON"); + assertThat(json.path("notes").get(0).path("timestamp_ms").asDouble()).isEqualTo(0.0); } } diff --git a/src/test/java/com/mr/global/client/ai/AiServerClientTest.java b/src/test/java/com/mr/global/client/ai/AiServerClientTest.java index ecd98299..c907f805 100644 --- a/src/test/java/com/mr/global/client/ai/AiServerClientTest.java +++ b/src/test/java/com/mr/global/client/ai/AiServerClientTest.java @@ -45,9 +45,9 @@ void setUp() { private AiAnalysisRequest sampleRequest() { return new AiAnalysisRequest( new AiAnalysisRequest.Meta(120.0, List.of(4, 4), - new AiAnalysisRequest.Key("C", "major"), "jazz"), + new AiAnalysisRequest.Key("C", "major"), "jazz", "basic"), List.of(new AiAnalysisRequest.Chord(1, 1.0, "Dm7")), - List.of(new AiAnalysisRequest.Note(0, 62, 0.0, 1.0, 90)) + List.of(new AiAnalysisRequest.Note(0, AiAnalysisRequest.NoteType.NOTE_ON, 62, 90, 0.0)) ); } diff --git a/src/test/java/com/mr/global/client/gemini/GeminiClientTest.java b/src/test/java/com/mr/global/client/gemini/GeminiClientTest.java new file mode 100644 index 00000000..ffe1f6da --- /dev/null +++ b/src/test/java/com/mr/global/client/gemini/GeminiClientTest.java @@ -0,0 +1,72 @@ +package com.mr.global.client.gemini; + +import static org.assertj.core.api.Assertions.assertThat; +import static org.springframework.test.web.client.match.MockRestRequestMatchers.header; +import static org.springframework.test.web.client.match.MockRestRequestMatchers.jsonPath; +import static org.springframework.test.web.client.match.MockRestRequestMatchers.requestTo; +import static org.springframework.test.web.client.response.MockRestResponseCreators.withSuccess; + +import com.mr.global.config.GeminiProperties; +import java.time.Duration; +import org.junit.jupiter.api.BeforeEach; +import org.junit.jupiter.api.Test; +import org.springframework.http.MediaType; +import org.springframework.test.web.client.MockRestServiceServer; +import org.springframework.web.client.RestClient; + +class GeminiClientTest { + + private static final String BASE_URL = "https://generativelanguage.googleapis.com"; + + private MockRestServiceServer server; + private GeminiClient client; + + @BeforeEach + void setUp() { + GeminiProperties properties = new GeminiProperties( + BASE_URL, + "test-key", + "gemini-3-flash-preview", + Duration.ofSeconds(5), + Duration.ofSeconds(60) + ); + RestClient.Builder builder = RestClient.builder().baseUrl(BASE_URL); + server = MockRestServiceServer.bindTo(builder).build(); + client = new GeminiClient(builder.build(), properties); + } + + @Test + void generateReport_callsConfiguredModelAndCombinesTextParts() { + server.expect(requestTo(BASE_URL + + "/v1beta/models/gemini-3-flash-preview:generateContent")) + .andExpect(header("x-goog-api-key", "test-key")) + .andExpect(jsonPath("$.generationConfig.maxOutputTokens").value(4096)) + .andRespond(withSuccess(""" + { + "candidates": [{ + "content": { + "parts": [ + {"text":"# 리포트\\n"}, + {"text":"내용"} + ] + } + }], + "usageMetadata": { + "promptTokenCount": 100, + "candidatesTokenCount": 50, + "totalTokenCount": 150, + "cachedContentTokenCount": 20 + } + } + """, MediaType.APPLICATION_JSON)); + + GeminiGenerationResult report = client.generateReport("system", "{}"); + + assertThat(report.content()).isEqualTo("# 리포트\n내용"); + assertThat(report.promptTokens()).isEqualTo(100); + assertThat(report.completionTokens()).isEqualTo(50); + assertThat(report.totalTokens()).isEqualTo(150); + assertThat(report.cacheHit()).isTrue(); + server.verify(); + } +}