Skip to content
Merged
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension


Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
4 changes: 4 additions & 0 deletions .claude/CLAUDE.md
Original file line number Diff line number Diff line change
@@ -1,5 +1,9 @@
# Studio Project Rules

## Branch Safety

- **Before making ANY changes**, always check that the user is NOT on the `main` or `stable` branch. Run `git branch --show-current` first. If on `main` or `stable`, STOP and tell the user to create/switch to a feature branch before proceeding.

## UI / Flutter

- Never use mobile-style animations or transitions in the Flutter app. This is a desktop application. Avoid `AnimatedContainer`, `AnimatedCrossFade`, `AnimatedRotation`, `AnimatedSwitcher`, `AnimatedOpacity`, `SlideTransition`, `FadeTransition`, swipe gestures, and similar animated widgets. Use instant state changes (e.g. `if`/`switch` conditionals, `Container`, `Transform.rotate`) instead.
Expand Down
106 changes: 106 additions & 0 deletions .github/scripts/pr_summary.py
Original file line number Diff line number Diff line change
@@ -0,0 +1,106 @@
"""Generate an AI summary for a pull request and update its description."""

import json
import os
import re
import subprocess
import urllib.request

ANTHROPIC_API_KEY = os.environ["ANTHROPIC_API_KEY"]
PR_NUMBER = os.environ["PR_NUMBER"]
BASE_REF = os.environ["BASE_REF"]

START_MARKER = "<!-- ai-summary-start -->"
END_MARKER = "<!-- ai-summary-end -->"

MAX_DIFF_CHARS = 80_000

PROMPT = (
"Summarize this pull request diff. Output a concise markdown summary with:\n"
"- A one-line overall summary\n"
"- A bullet list of key changes grouped by area\n\n"
"Keep it short and useful for code reviewers. "
"Do not include any preamble, just output the summary.\n\n"
"Diff:\n"
)


def get_diff() -> str:
result = subprocess.run(
["git", "diff", f"origin/{BASE_REF}...HEAD", "--", ".", ":!*.lock", ":!*.sum"],
capture_output=True,
text=True,
check=True,
)
return result.stdout[:MAX_DIFF_CHARS]


def call_claude(diff: str) -> str:
payload = json.dumps({
"model": "claude-haiku-4-5-20251001",
"max_tokens": 1024,
"messages": [{"role": "user", "content": PROMPT + diff}],
}).encode()

req = urllib.request.Request(
"https://api.anthropic.com/v1/messages",
data=payload,
headers={
"Content-Type": "application/json",
"X-Api-Key": ANTHROPIC_API_KEY,
"Anthropic-Version": "2023-06-01",
},
)

with urllib.request.urlopen(req) as resp:
body = json.loads(resp.read())

return body["content"][0]["text"]


def get_pr_body() -> str:
result = subprocess.run(
["gh", "pr", "view", PR_NUMBER, "--json", "body", "-q", ".body"],
capture_output=True,
text=True,
check=True,
)
return result.stdout.strip()


def update_pr_body(new_body: str) -> None:
subprocess.run(
["gh", "pr", "edit", PR_NUMBER, "--body", new_body],
check=True,
)


def main() -> None:
diff = get_diff()
if not diff.strip():
print("No diff found, skipping summary.")
return

print("Calling Claude API...")
summary = call_claude(diff)

ai_section = f"{START_MARKER}\n## Summary (AI-generated)\n{summary}\n{END_MARKER}"

current_body = get_pr_body()

pattern = re.compile(
re.escape(START_MARKER) + r".*?" + re.escape(END_MARKER),
re.DOTALL,
)

if pattern.search(current_body):
new_body = pattern.sub(ai_section, current_body)
else:
new_body = f"{ai_section}\n\n{current_body}" if current_body else ai_section

update_pr_body(new_body)
print("PR description updated with AI summary.")


if __name__ == "__main__":
main()
25 changes: 25 additions & 0 deletions .github/workflows/pr-summary.yml
Original file line number Diff line number Diff line change
@@ -0,0 +1,25 @@
name: PR Summary

on:
pull_request:
types: [opened, synchronize]

permissions:
pull-requests: write
contents: read

jobs:
summarize:
runs-on: ubuntu-latest
steps:
- uses: actions/checkout@v4
with:
fetch-depth: 0

- name: Generate summary
env:
ANTHROPIC_API_KEY: ${{ secrets.ANTHROPIC_API_KEY }}
GH_TOKEN: ${{ secrets.GITHUB_TOKEN }}
PR_NUMBER: ${{ github.event.pull_request.number }}
BASE_REF: ${{ github.event.pull_request.base.ref }}
run: python3 .github/scripts/pr_summary.py
Original file line number Diff line number Diff line change
Expand Up @@ -2,6 +2,7 @@ import 'dart:convert';

import 'package:drift/drift.dart';
import 'package:studio_backend/src/audio/dto/audio_generate_request.dart';
import 'package:studio_backend/src/audio/task_status_values.dart';
import 'package:studio_backend/src/database/database.dart';
import 'package:studio_backend/src/database/postgres.dart';
import 'package:studio_backend/src/utils/cursor_pagination.dart';
Expand Down Expand Up @@ -38,7 +39,7 @@ class AudioGenerationTaskRepositoryImpl
taskId: Value(taskId),
model: Value(request.model),
taskType: Value(request.taskType),
status: const Value('processing'),
status: const Value(TaskStatusValues.processing),
prompt: Value(request.prompt),
lyrics: Value(request.lyrics),
negativePrompt: Value(request.negativePrompt),
Expand Down Expand Up @@ -85,7 +86,7 @@ class AudioGenerationTaskRepositoryImpl
taskId: taskId,
model: request.model,
taskType: request.taskType,
status: 'processing',
status: TaskStatusValues.processing,
title: Value(request.title),
prompt: Value(request.prompt),
lyrics: Value(request.lyrics),
Expand Down Expand Up @@ -132,7 +133,7 @@ class AudioGenerationTaskRepositoryImpl
_database.audioGenerationTask,
)..where((t) => t.taskId.equals(taskId))).write(
AudioGenerationTaskCompanion(
status: const Value('complete'),
status: const Value(TaskStatusValues.complete),
result: Value(jsonEncode(result)),
completedAt: Value(DateTime.now().toPgDateTime()),
),
Expand All @@ -148,7 +149,7 @@ class AudioGenerationTaskRepositoryImpl
_database.audioGenerationTask,
)..where((t) => t.taskId.equals(taskId))).write(
AudioGenerationTaskCompanion(
status: const Value('failed'),
status: const Value(TaskStatusValues.failed),
error: Value(error),
completedAt: Value(DateTime.now().toPgDateTime()),
),
Expand All @@ -169,7 +170,7 @@ class AudioGenerationTaskRepositoryImpl
..where(
(t) {
var clause = t.userId.equals(userId) &
t.status.equals('complete') &
t.status.equals(TaskStatusValues.complete) &
buildCursorWhereClause(t.createdAt, t.id, cursor,
descending: descending);
if (rating != null) {
Expand Down Expand Up @@ -207,7 +208,7 @@ class AudioGenerationTaskRepositoryImpl
taskId: taskId,
model: 'upload',
taskType: 'upload',
status: 'uploading',
status: TaskStatusValues.uploading,
srcAudioPath: Value(objectPath),
workspaceId: Value(workspaceId),
),
Expand Down Expand Up @@ -298,7 +299,7 @@ class AudioGenerationTaskRepositoryImpl
..where((t) =>
t.lyricSheetId.equals(lyricSheetId) &
t.userId.equals(userId) &
t.status.equals('complete'))
t.status.equals(TaskStatusValues.complete))
..orderBy([(t) => OrderingTerm.desc(t.createdAt)]))
.get();
}
Expand Down Expand Up @@ -337,7 +338,7 @@ class AudioGenerationTaskRepositoryImpl
final q = _database.select(_database.audioGenerationTask)
..where((t) {
var clause = t.userId.equals(userId) &
t.status.equals('complete') &
t.status.equals(TaskStatusValues.complete) &
t.lyrics.lower().like(pattern.toLowerCase()) &
buildCursorWhereClause(t.createdAt, t.id, cursor,
descending: descending);
Expand Down
11 changes: 6 additions & 5 deletions packages/studio_backend/lib/src/audio/audio_service.dart
Original file line number Diff line number Diff line change
Expand Up @@ -4,6 +4,7 @@ import 'dart:typed_data';

import 'package:http/http.dart' as http;
import 'package:studio_backend/src/audio/audio_client.dart';
import 'package:studio_backend/src/audio/task_status_values.dart';
import 'package:studio_backend/src/audio/audio_generation_task_repository.dart';
import 'package:studio_backend/src/audio/audio_model_client.dart';
import 'package:studio_backend/src/audio/dto/audio_generate_request.dart';
Expand Down Expand Up @@ -695,7 +696,7 @@ class AudioService {
await _setTask(
_TaskStatus(
taskId: taskId,
status: 'processing',
status: TaskStatusValues.processing,
taskType: generateRequest.taskType,
model: generateRequest.model,
),
Expand Down Expand Up @@ -740,10 +741,10 @@ class AudioService {
if (task.model != null) 'model': task.model,
};

if (task.status == 'complete' && task.result != null) {
if (task.status == TaskStatusValues.complete && task.result != null) {
response['result'] = task.result;
}
if (task.status == 'failed' && task.error != null) {
if (task.status == TaskStatusValues.failed && task.error != null) {
response['error'] = task.error;
}

Expand All @@ -765,7 +766,7 @@ class AudioService {
await _setTask(
_TaskStatus(
taskId: taskId,
status: 'complete',
status: TaskStatusValues.complete,
taskType: prev?.taskType ?? 'unknown',
model: prev?.model,
result: result,
Expand All @@ -787,7 +788,7 @@ class AudioService {
await _setTask(
_TaskStatus(
taskId: taskId,
status: 'failed',
status: TaskStatusValues.failed,
taskType: prev?.taskType ?? 'unknown',
model: prev?.model,
error: 'Generation failed',
Expand Down
5 changes: 3 additions & 2 deletions packages/studio_backend/lib/src/audio/midi/midi_client.dart
Original file line number Diff line number Diff line change
@@ -1,4 +1,5 @@
import 'package:studio_backend/src/audio/audio_model_client.dart';
import 'package:studio_backend/src/audio/task_status_values.dart';

/// Anticorruption layer for the MIDI generation model.
///
Expand Down Expand Up @@ -94,7 +95,7 @@ class MidiClient extends AudioModelClient {
final response = await getRequest('/api/tasks/$taskId/');
final status = response['status'] as String?;

if (status == 'complete') {
if (status == TaskStatusValues.complete) {
final results = <Map<String, dynamic>>[];
final downloadUrl = response['download_url'] as String?;
final mp3DownloadUrl = response['mp3_download_url'] as String?;
Expand All @@ -111,7 +112,7 @@ class MidiClient extends AudioModelClient {
return {'results': results};
}

if (status == 'failed') {
if (status == TaskStatusValues.failed) {
throw AudioModelException(
500,
'MIDI task $taskId failed: '
Expand Down
7 changes: 7 additions & 0 deletions packages/studio_backend/lib/src/audio/task_status_values.dart
Original file line number Diff line number Diff line change
@@ -0,0 +1,7 @@
/// Canonical status strings for audio generation tasks.
abstract final class TaskStatusValues {
static const processing = 'processing';
static const uploading = 'uploading';
static const complete = 'complete';
static const failed = 'failed';
}
6 changes: 0 additions & 6 deletions packages/studio_ui/lib/configuration/configuration_base.dart
Original file line number Diff line number Diff line change
Expand Up @@ -11,12 +11,6 @@ class Configuration {
final bool secure;
final String applicationId;

Uri buildUri(String path, [Map<String, dynamic>? query]) {
return secure
? Uri.https(apiHost, path, query)
: Uri.http(apiHost, path, query);
}

static String environmentLookup() {
const envFromDefine = String.fromEnvironment('BUILD_ENV');
if (envFromDefine.isNotEmpty) return envFromDefine;
Expand Down
15 changes: 10 additions & 5 deletions packages/studio_ui/lib/models/task_status.dart
Original file line number Diff line number Diff line change
@@ -1,4 +1,9 @@
class TaskStatus {
static const statusProcessing = 'processing';
static const statusUploading = 'uploading';
static const statusComplete = 'complete';
static const statusFailed = 'failed';

TaskStatus({
required this.taskId,
required this.status,
Expand Down Expand Up @@ -28,7 +33,7 @@ class TaskStatus {

factory TaskStatus.fromSongJson(Map<String, dynamic> json) => TaskStatus(
taskId: json['task_id'] as String,
status: json['status'] as String? ?? 'complete',
status: json['status'] as String? ?? statusComplete,
taskType: json['task_type'] as String? ?? 'unknown',
result: json['result'] as Map<String, dynamic>?,
prompt: json['prompt'] as String?,
Expand Down Expand Up @@ -57,10 +62,10 @@ class TaskStatus {
final DateTime? createdAt;
final Map<String, dynamic>? parameters;

bool get isProcessing => status == 'processing';
bool get isUploading => status == 'uploading';
bool get isComplete => status == 'complete';
bool get isFailed => status == 'failed';
bool get isProcessing => status == statusProcessing;
bool get isUploading => status == statusUploading;
bool get isComplete => status == statusComplete;
bool get isFailed => status == statusFailed;

/// Whether this task is still active and should be polled.
bool get isActive => isProcessing || isUploading;
Expand Down
14 changes: 7 additions & 7 deletions packages/studio_ui/lib/pages/create_page.dart
Original file line number Diff line number Diff line change
Expand Up @@ -952,7 +952,7 @@ class _CreatePageState extends State<CreatePage> {
final taskId = await client.submitTask(body);
final status = TaskStatus(
taskId: taskId,
status: 'processing',
status: TaskStatus.statusProcessing,
taskType: _taskType,
model: _model,
);
Expand Down Expand Up @@ -1008,7 +1008,7 @@ class _CreatePageState extends State<CreatePage> {
setState(() {
_tasks.insert(
0,
TaskStatus(taskId: fileId, status: 'uploading', taskType: 'upload'),
TaskStatus(taskId: fileId, status: TaskStatus.statusUploading, taskType: 'upload'),
);
_uploadProgress = 0.2;
});
Expand All @@ -1034,7 +1034,7 @@ class _CreatePageState extends State<CreatePage> {
setState(() {
_tasks[index] = TaskStatus(
taskId: fileId,
status: 'complete',
status: TaskStatus.statusComplete,
taskType: 'upload',
);
_pickedFile = null;
Expand Down Expand Up @@ -5146,10 +5146,10 @@ class _TaskCardState extends State<_TaskCard>
Widget _statusBadge() {
final s = S.of(context);
final (color, label) = switch (widget.task.status) {
'processing' => (AppColors.controlPink, s.statusProcessing),
'uploading' => (Colors.orange, s.statusUploading),
'complete' => (Colors.green, s.statusComplete),
'failed' => (Colors.redAccent, s.statusFailed),
TaskStatus.statusProcessing => (AppColors.controlPink, s.statusProcessing),
TaskStatus.statusUploading => (Colors.orange, s.statusUploading),
TaskStatus.statusComplete => (Colors.green, s.statusComplete),
TaskStatus.statusFailed => (Colors.redAccent, s.statusFailed),
_ => (AppColors.textMuted, widget.task.status),
};

Expand Down
Loading