diff --git a/.github/workflows/test-postgres.yml b/.github/workflows/test-postgres.yml new file mode 100644 index 0000000000..09aa883855 --- /dev/null +++ b/.github/workflows/test-postgres.yml @@ -0,0 +1,59 @@ +name: Test Postgres + +on: + push: + branches: [ postgres-support-3 ] + pull_request: + release: + types: [ published ] + +jobs: + test-postgres: + runs-on: ubuntu-24.04 + services: + postgres: + image: postgres:17 + env: + POSTGRES_USER: username + POSTGRES_PASSWORD: password + POSTGRES_DB: testdb + ports: + - 5432:5432 + options: >- + --name postgres-db + --health-cmd pg_isready + --health-interval 10s + --health-timeout 5s + --health-retries 5 + steps: + - uses: actions/checkout@v2 + + - name: Setup Go + uses: actions/setup-go@v5 + with: + go-version-file: 'go.mod' + + - name: Setup Node.js + uses: actions/setup-node@v4 + with: + node-version: '20' + + - name: Setup pnpm + uses: pnpm/action-setup@v4 + with: + version: 9 + + - name: Build project + run: | + rm -f ui/v2.5/pnpm-workspace.yaml + make release + + - name: Install plperl in postgres container + run: | + docker exec postgres-db apt-get update + docker exec postgres-db apt-get install -y postgresql-plperl-17 + + - name: Test + env: + PGSQL_TEST: 'postgresql://username:password@localhost:5432/testdb' + run: go test -tags "db_integration" ./... diff --git a/go.mod b/go.mod index db0d6fe34e..c994c22500 100644 --- a/go.mod +++ b/go.mod @@ -14,7 +14,7 @@ require ( github.com/corona10/goimagehash v1.1.0 github.com/disintegration/imaging v1.6.2 github.com/dop251/goja v0.0.0-20231027120936-b396bb4c349d - github.com/doug-martin/goqu/v9 v9.18.0 + github.com/doug-martin/goqu/v9 v9.19.1-0.20231214054827-21b6e6d1cb1b github.com/go-chi/chi/v5 v5.2.2 github.com/go-chi/cors v1.2.1 github.com/go-chi/httplog v0.3.1 @@ -27,6 +27,7 @@ require ( github.com/gorilla/websocket v1.5.0 github.com/hashicorp/golang-lru/v2 v2.0.7 github.com/hasura/go-graphql-client v0.13.1 + github.com/jackc/pgx/v5 v5.7.6 github.com/jinzhu/copier v0.4.0 github.com/jmoiron/sqlx v1.4.0 github.com/json-iterator/go v1.1.12 @@ -91,7 +92,11 @@ require ( github.com/hashicorp/go-multierror v1.1.1 // indirect github.com/hashicorp/hcl v1.0.0 // indirect github.com/inconshreveable/mousetrap v1.1.0 // indirect + github.com/jackc/pgpassfile v1.0.0 // indirect + github.com/jackc/pgservicefile v0.0.0-20240606120523-5a60cdf6a761 // indirect + github.com/jackc/puddle/v2 v2.2.2 // indirect github.com/knadh/koanf/maps v0.1.2 // indirect + github.com/lib/pq v1.10.9 // indirect github.com/magiconair/properties v1.8.7 // indirect github.com/mattn/go-colorable v0.1.14 // indirect github.com/mattn/go-isatty v0.0.20 // indirect diff --git a/go.sum b/go.sum index dbe82cf99f..a198ff86b8 100644 --- a/go.sum +++ b/go.sum @@ -53,11 +53,15 @@ filippo.io/edwards25519 v1.1.0 h1:FNf4tywRC1HmFuKW5xopWpigGjJKiJSV0Cqo0cJWDaA= filippo.io/edwards25519 v1.1.0/go.mod h1:BxyFTGdWcka3PhytdK4V28tE5sGfRvvvRV7EaN4VDT4= github.com/99designs/gqlgen v0.17.73 h1:A3Ki+rHWqKbAOlg5fxiZBnz6OjW3nwupDHEG15gEsrg= github.com/99designs/gqlgen v0.17.73/go.mod h1:2RyGWjy2k7W9jxrs8MOQthXGkD3L3oGr0jXW3Pu8lGg= +github.com/Azure/go-ansiterm v0.0.0-20230124172434-306776ec8161 h1:L/gRVlceqvL25UVaW/CKtUDjefjrs0SPonmDGUVOYP0= +github.com/Azure/go-ansiterm v0.0.0-20230124172434-306776ec8161/go.mod h1:xomTg63KZ2rFqZQzSB4Vz2SUXa1BpHTVz9L5PTmPC4E= github.com/BurntSushi/toml v0.3.1/go.mod h1:xHWCNGjB5oqiDr8zfno3MHue2Ht5sIBksp03qcyfWMU= github.com/BurntSushi/xgb v0.0.0-20160522181843-27f122750802/go.mod h1:IVnqGOEym/WlBOVXweHU+Q+/VP0lqqI8lqeDx9IjBqo= github.com/DATA-DOG/go-sqlmock v1.5.0 h1:Shsta01QNfFxHCfpW6YH2STWB0MudeXXEWMr20OEh60= github.com/DATA-DOG/go-sqlmock v1.5.0/go.mod h1:f/Ixk793poVmq4qj/V1dPUg2JEAKC73Q5eFN3EC/SaM= github.com/DataDog/datadog-go v3.2.0+incompatible/go.mod h1:LButxg5PwREeZtORoXG3tL4fMGNddJ+vMq1mwgfaqoQ= +github.com/Microsoft/go-winio v0.6.1 h1:9/kr64B9VUZrLm5YYwbGtUJnMgqWVOdUAXu6Migciow= +github.com/Microsoft/go-winio v0.6.1/go.mod h1:LRdKpFKfdobln8UmuiYcKPot9D2v6svN5+sAH+4kjUM= github.com/OneOfOne/xxhash v1.2.2/go.mod h1:HSdplMjZKSmBqAxg5vPj2TmRDmfkzw+cTzAElWljhcU= github.com/PuerkitoBio/goquery v1.10.3 h1:pFYcNSqHxBD06Fpj/KsbStFRsgRATgnf3LeXiUkhzPo= github.com/PuerkitoBio/goquery v1.10.3/go.mod h1:tMUX0zDMHXYlAQk6p35XxQMqMweEKB7iK7iLNd4RH4Y= @@ -156,22 +160,31 @@ github.com/creack/pty v1.1.9/go.mod h1:oKZEueFk5CKHvIhNR5MUki03XCEU+Q6VDXinZuGJ3 github.com/davecgh/go-spew v1.1.0/go.mod h1:J7Y8YcW2NihsgmVo/mv3lAwl/skON4iLHjSsI+c5H38= github.com/davecgh/go-spew v1.1.1 h1:vj9j/u1bqnvCEfJOwUhtlOARqs3+rkHYY13jYWTU97c= github.com/davecgh/go-spew v1.1.1/go.mod h1:J7Y8YcW2NihsgmVo/mv3lAwl/skON4iLHjSsI+c5H38= -github.com/denisenkom/go-mssqldb v0.10.0/go.mod h1:xbL0rPBG9cCiLr28tMa8zpbdarY27NDyej4t/EjAShU= github.com/dgryski/trifles v0.0.0-20230903005119-f50d829f2e54 h1:SG7nF6SRlWhcT7cNTs5R6Hk4V2lcmLz2NsG2VnInyNo= github.com/dgryski/trifles v0.0.0-20230903005119-f50d829f2e54/go.mod h1:if7Fbed8SFyPtHLHbg49SI7NAdJiC5WIA09pe59rfAA= +github.com/dhui/dktest v0.3.16 h1:i6gq2YQEtcrjKbeJpBkWjE8MmLZPYllcjOFbTZuPDnw= +github.com/dhui/dktest v0.3.16/go.mod h1:gYaA3LRmM8Z4vJl2MA0THIigJoZrwOansEOsp+kqxp0= github.com/disintegration/imaging v1.6.2 h1:w1LecBlG2Lnp8B3jk5zSuNqd7b4DXhcjwek1ei82L+c= github.com/disintegration/imaging v1.6.2/go.mod h1:44/5580QXChDfwIclfc/PCwrr44amcmDAg8hxG0Ewe4= github.com/dlclark/regexp2 v1.4.1-0.20201116162257-a2a8dda75c91/go.mod h1:2pZnwuY/m+8K6iRw6wQdMtk+rH5tNGR1i55kozfMjCc= github.com/dlclark/regexp2 v1.7.0 h1:7lJfhqlPssTb1WQx4yvTHN0uElPEv52sbaECrAQxjAo= github.com/dlclark/regexp2 v1.7.0/go.mod h1:DHkYz0B9wPfa6wondMfaivmHpzrQ3v9q8cnmRbL6yW8= +github.com/docker/distribution v2.8.2+incompatible h1:T3de5rq0dB1j30rp0sA2rER+m322EBzniBPB6ZIzuh8= +github.com/docker/distribution v2.8.2+incompatible/go.mod h1:J2gT2udsDAN96Uj4KfcMRqY0/ypR+oyYUYmja8H+y+w= +github.com/docker/docker v20.10.24+incompatible h1:Ugvxm7a8+Gz6vqQYQQ2W7GYq5EUPaAiuPgIfVyI3dYE= +github.com/docker/docker v20.10.24+incompatible/go.mod h1:eEKB0N0r5NX/I1kEveEz05bcu8tLC/8azJZsviup8Sk= +github.com/docker/go-connections v0.4.0 h1:El9xVISelRB7BuFusrZozjnkIM5YnzCViNKohAFqRJQ= +github.com/docker/go-connections v0.4.0/go.mod h1:Gbd7IOopHjR8Iph03tsViu4nIes5XhDvyHbTtUxmeec= +github.com/docker/go-units v0.5.0 h1:69rxXcBk27SvSaaxTtLh/8llcHD8vYHT7WSdRZ/jvr4= +github.com/docker/go-units v0.5.0/go.mod h1:fgPhTUdO+D/Jk86RDLlptpiXQzgHJF7gydDDbaIK4Dk= github.com/docopt/docopt-go v0.0.0-20180111231733-ee0de3bc6815/go.mod h1:WwZ+bS3ebgob9U8Nd0kOddGdZWjyMGR8Wziv+TBNwSE= github.com/dop251/goja v0.0.0-20211022113120-dc8c55024d06/go.mod h1:R9ET47fwRVRPZnOGvHxxhuZcbrMCuiqOz3Rlrh4KSnk= github.com/dop251/goja v0.0.0-20231027120936-b396bb4c349d h1:wi6jN5LVt/ljaBG4ue79Ekzb12QfJ52L9Q98tl8SWhw= github.com/dop251/goja v0.0.0-20231027120936-b396bb4c349d/go.mod h1:QMWlm50DNe14hD7t24KEqZuUdC9sOTy8W6XbCU1mlw4= github.com/dop251/goja_nodejs v0.0.0-20210225215109-d91c329300e7/go.mod h1:hn7BA7c8pLvoGndExHudxTDKZ84Pyvv+90pbBjbTz0Y= github.com/dop251/goja_nodejs v0.0.0-20211022123610-8dd9abb0616d/go.mod h1:DngW8aVqWbuLRMHItjPUyqdj+HWPvnQe8V8y1nDpIbM= -github.com/doug-martin/goqu/v9 v9.18.0 h1:/6bcuEtAe6nsSMVK/M+fOiXUNfyFF3yYtE07DBPFMYY= -github.com/doug-martin/goqu/v9 v9.18.0/go.mod h1:nf0Wc2/hV3gYK9LiyqIrzBEVGlI8qW3GuDCEobC4wBQ= +github.com/doug-martin/goqu/v9 v9.19.1-0.20231214054827-21b6e6d1cb1b h1:WaCes6lOJCbIDgABfA8gB1ADMQo6+ftGEkj+oIB+vm4= +github.com/doug-martin/goqu/v9 v9.19.1-0.20231214054827-21b6e6d1cb1b/go.mod h1:1MqhYk2p5QFEUT9ZzH+M02Jv8BbOYlvzupULdHl7Mjs= github.com/dustin/go-humanize v0.0.0-20180421182945-02af3965c54e/go.mod h1:HtrtbFcZ19U5GC7JDqmcUSB87Iq5E25KnS6fMYU6eOk= github.com/envoyproxy/go-control-plane v0.9.0/go.mod h1:YTl/9mNaCwkRvm6d1a2C3ymFceY/DCBVvsKhRF0iEA4= github.com/envoyproxy/go-control-plane v0.9.1-0.20191026205805-5f8ba28d4473/go.mod h1:YTl/9mNaCwkRvm6d1a2C3ymFceY/DCBVvsKhRF0iEA4= @@ -213,7 +226,6 @@ github.com/go-logfmt/logfmt v0.3.0/go.mod h1:Qt1PoO58o5twSAckw1HlFXLmHsOX5/0LbT9 github.com/go-logfmt/logfmt v0.4.0/go.mod h1:3RMwSq7FuexP4Kalkev3ejPJsZTpXXBr9+V4qmtdjCk= github.com/go-sourcemap/sourcemap v2.1.3+incompatible h1:W1iEw64niKVGogNgBN3ePyLFfuisuzeidWPMPWmECqU= github.com/go-sourcemap/sourcemap v2.1.3+incompatible/go.mod h1:F8jJfvm2KbVjc5NqelyYJmf/v5J0dwNLS2mL4sNA1Jg= -github.com/go-sql-driver/mysql v1.6.0/go.mod h1:DCzpHaOWr8IXmIStZouvnhqoel9Qv2LBy8hT2VhHyBg= github.com/go-sql-driver/mysql v1.8.1 h1:LedoTUt/eveggdHS9qUFC1EFSa8bU2+1pZjSRpvNJ1Y= github.com/go-sql-driver/mysql v1.8.1/go.mod h1:wEBSXgmK//2ZFJyE+qWnIsVGmvmEKlqwuVSjsCm7DZg= github.com/go-stack/stack v1.8.0/go.mod h1:v0f6uXyyMGvRgIKkXu+yp6POWl0qKG85gN/melR3HDY= @@ -231,12 +243,12 @@ github.com/goccy/go-yaml v1.18.0 h1:8W7wMFS12Pcas7KU+VVkaiCng+kG8QiFeFwzFb+rwuw= github.com/goccy/go-yaml v1.18.0/go.mod h1:XBurs7gK8ATbW4ZPGKgcbrY1Br56PdM69F7LkFRi1kA= github.com/godbus/dbus/v5 v5.0.4/go.mod h1:xhWf0FNVPg57R7Z0UbKHbJfkEywrmjJnf7w5xrFpKfA= github.com/gogo/protobuf v1.1.1/go.mod h1:r8qH/GZQm5c6nD/R0oafs1akxWv10x8SbQlK7atdtwQ= +github.com/gogo/protobuf v1.3.2 h1:Ov1cvc58UF3b5XjBnZv7+opcTcQFZebYjWzi34vdm4Q= github.com/gogo/protobuf v1.3.2/go.mod h1:P1XiOD3dCwIKUDQYPy72D8LYyHL2YPYrpS2s69NZV8Q= github.com/golang-jwt/jwt/v4 v4.5.2 h1:YtQM7lnr8iZ+j5q71MGKkNw9Mn7AjHM68uc9g5fXeUI= github.com/golang-jwt/jwt/v4 v4.5.2/go.mod h1:m21LjoU+eqJr34lmDMbreY2eSTRJ1cv77w39/MY0Ch0= github.com/golang-migrate/migrate/v4 v4.16.2 h1:8coYbMKUyInrFk1lfGfRovTLAW7PhWp8qQDT2iKfuoA= github.com/golang-migrate/migrate/v4 v4.16.2/go.mod h1:pfcJX4nPHaVdc5nmdCikFBWtm+UBpiZjRNNsyBbp0/o= -github.com/golang-sql/civil v0.0.0-20190719163853-cb61b32ac6fe/go.mod h1:8vg3r2VgvsThLBIFL93Qb5yWzgyZWhEmBwUJWevAkK0= github.com/golang/glog v0.0.0-20160126235308-23def4e6c14b/go.mod h1:SBH7ygxi8pfUlaOkMMuAQtPIUF8ecWP5IEl/CR7VP2Q= github.com/golang/groupcache v0.0.0-20190702054246-869f871628b6/go.mod h1:cIg4eruTrX1D+g88fzRXU5OdNfaM+9IcxsU14FzY7Hc= github.com/golang/groupcache v0.0.0-20191227052852-215e87163ea7/go.mod h1:cIg4eruTrX1D+g88fzRXU5OdNfaM+9IcxsU14FzY7Hc= @@ -376,6 +388,14 @@ github.com/ianlancetaylor/demangle v0.0.0-20220319035150-800ac71e25c2/go.mod h1: github.com/inconshreveable/mousetrap v1.0.0/go.mod h1:PxqpIevigyE2G7u3NXJIT2ANytuPF1OarO4DADm73n8= github.com/inconshreveable/mousetrap v1.1.0 h1:wN+x4NVGpMsO7ErUn/mUI3vEoE6Jt13X2s0bqwp9tc8= github.com/inconshreveable/mousetrap v1.1.0/go.mod h1:vpF70FUmC8bwa3OWnCshd2FqLfsEA9PFc4w1p2J65bw= +github.com/jackc/pgpassfile v1.0.0 h1:/6Hmqy13Ss2zCq62VdNG8tM1wchn8zjSGOBJ6icpsIM= +github.com/jackc/pgpassfile v1.0.0/go.mod h1:CEx0iS5ambNFdcRtxPj5JhEz+xB6uRky5eyVu/W2HEg= +github.com/jackc/pgservicefile v0.0.0-20240606120523-5a60cdf6a761 h1:iCEnooe7UlwOQYpKFhBabPMi4aNAfoODPEFNiAnClxo= +github.com/jackc/pgservicefile v0.0.0-20240606120523-5a60cdf6a761/go.mod h1:5TJZWKEWniPve33vlWYSoGYefn3gLQRzjfDlhSJ9ZKM= +github.com/jackc/pgx/v5 v5.7.6 h1:rWQc5FwZSPX58r1OQmkuaNicxdmExaEz5A2DO2hUuTk= +github.com/jackc/pgx/v5 v5.7.6/go.mod h1:aruU7o91Tc2q2cFp5h4uP3f6ztExVpyVv88Xl/8Vl8M= +github.com/jackc/puddle/v2 v2.2.2 h1:PR8nw+E/1w0GLuRFSmiioY6UooMp6KJv0/61nB7icHo= +github.com/jackc/puddle/v2 v2.2.2/go.mod h1:vriiEXHvEE654aYKXXjOvZM39qJ0q+azkZFrfEOc3H4= github.com/jinzhu/copier v0.4.0 h1:w3ciUoD19shMCRargcpm0cm91ytaBhDvuRpz1ODO/U8= github.com/jinzhu/copier v0.4.0/go.mod h1:DfbEm0FYsaqBcKcFuvmOZb218JkPGtvSHsKg8S8hyyg= github.com/jmoiron/sqlx v1.4.0 h1:1PLqN7S1UYp5t4SrVVnt4nUVNemrDAtxlulVe+Qgm3o= @@ -422,7 +442,6 @@ github.com/kr/text v0.2.0 h1:5Nx0Ya0ZqY2ygV366QzturHI13Jq95ApcVaJBhpS+AY= github.com/kr/text v0.2.0/go.mod h1:eLer722TekiGuMkidMxC/pM04lWEeraHUUmBw8l2grE= github.com/ledongthuc/pdf v0.0.0-20220302134840-0c2507a12d80 h1:6Yzfa6GP0rIo/kULo2bwGEkFvCePZ3qHDDTC3/J9Swo= github.com/ledongthuc/pdf v0.0.0-20220302134840-0c2507a12d80/go.mod h1:imJHygn/1yfhB7XSJJKlFZKl/J+dCPAknuiaGOshXAs= -github.com/lib/pq v1.10.1/go.mod h1:AlVN5x4E4T544tWzH6hKfbfQvm3HdbOxrmggDNAPY9o= github.com/lib/pq v1.10.9 h1:YXG7RB+JIjhP29X+OtkiDnYaXQwpS4JEWq7dtCCRUEw= github.com/lib/pq v1.10.9/go.mod h1:AlVN5x4E4T544tWzH6hKfbfQvm3HdbOxrmggDNAPY9o= github.com/lucasb-eyer/go-colorful v1.2.0 h1:1nnpGOrhyZZuNyfu1QjKiUICQ74+3FNCN69Aj6K7nkY= @@ -449,7 +468,6 @@ github.com/mattn/go-isatty v0.0.16/go.mod h1:kYGgaQfpe5nmfYZH+SKPsOc2e4SrIfOl2e/ github.com/mattn/go-isatty v0.0.19/go.mod h1:W+V8PltTTMOvKvAeJH7IuucS94S2C6jfK/D7dTCTo3Y= github.com/mattn/go-isatty v0.0.20 h1:xfD0iDuEKnDkl03q4limB+vH+GxLEtL/jb4xVJSWWEY= github.com/mattn/go-isatty v0.0.20/go.mod h1:W+V8PltTTMOvKvAeJH7IuucS94S2C6jfK/D7dTCTo3Y= -github.com/mattn/go-sqlite3 v1.14.7/go.mod h1:NyWgC/yNuGj7Q9rpYnZvas74GogHl5/Z4A/KQRfk6bU= github.com/mattn/go-sqlite3 v1.14.22 h1:2gZY6PC6kBnID23Tichd1K+Z0oS6nE/XwU+Vz/5o4kU= github.com/mattn/go-sqlite3 v1.14.22/go.mod h1:Uh1q+B4BYcTPb+yiD3kU8Ct7aC0hY9fxUwlHK0RXw+Y= github.com/matttproud/golang_protobuf_extensions v1.0.1/go.mod h1:D8He9yQNgCq6Z5Ld7szi9bcBfOoFv/3dc6xSMkL2PC0= @@ -469,6 +487,8 @@ github.com/mitchellh/mapstructure v1.5.0 h1:jeMsZIYE/09sWLaz43PL7Gy6RuMjD2eJVyua github.com/mitchellh/mapstructure v1.5.0/go.mod h1:bFUtVrKA4DC2yAKiSyO/QUcy7e+RRV2QTWOzhPopBRo= github.com/mitchellh/reflectwalk v1.0.2 h1:G2LzWKi524PWgd3mLHV8Y5k7s6XUvT0Gef6zxSIeXaQ= github.com/mitchellh/reflectwalk v1.0.2/go.mod h1:mSTlrgnPZtwu0c4WaC2kGObEpuNDbx0jmZXqmk4esnw= +github.com/moby/term v0.5.0 h1:xt8Q1nalod/v7BqbG21f8mQPqH+xAaC9C3N3wfWbVP0= +github.com/moby/term v0.5.0/go.mod h1:8FzsFHVUBGZdbDsJw/ot+X+d5HLUbvklYLJ9uGfcI3Y= github.com/modern-go/concurrent v0.0.0-20180228061459-e0a39a4cb421/go.mod h1:6dJC0mAP4ikYIbvyc7fijjWJddQyLn8Ig3JB5CqoB9Q= github.com/modern-go/concurrent v0.0.0-20180306012644-bacd9c7ef1dd h1:TRLaZ9cD/w8PVh93nsPXa1VrQ6jlwL5oN8l14QlcNfg= github.com/modern-go/concurrent v0.0.0-20180306012644-bacd9c7ef1dd/go.mod h1:6dJC0mAP4ikYIbvyc7fijjWJddQyLn8Ig3JB5CqoB9Q= @@ -476,6 +496,8 @@ github.com/modern-go/reflect2 v0.0.0-20180701023420-4b7aa43c6742/go.mod h1:bx2lN github.com/modern-go/reflect2 v1.0.1/go.mod h1:bx2lNnkwVCuqBIxFjflWJWanXIb3RllmbCylyMrvgv0= github.com/modern-go/reflect2 v1.0.2 h1:xBagoLtFs94CBntxluKeaWgTMpvLxC4ur3nMaC9Gz0M= github.com/modern-go/reflect2 v1.0.2/go.mod h1:yWuevngMOJpCy52FWWMvUC8ws7m/LJsjYzDa0/r8luk= +github.com/morikuni/aec v1.0.0 h1:nP9CBfwrvYnBRgY6qfDQkygYDmYwOilePFkwzv4dU8A= +github.com/morikuni/aec v1.0.0/go.mod h1:BbKIizmSmc5MMPqRYbxO4ZU0S0+P200+tUnFx7PXmsc= github.com/mschoch/smat v0.0.0-20160514031455-90eadee771ae/go.mod h1:qAyveg+e4CE+eKJXWVjKXM4ck2QobLqTDytGJbLLhJg= github.com/mwitkow/go-conntrack v0.0.0-20161129095857-cc309e4a2223/go.mod h1:qRWi+5nqEBWmkhHvq77mSJWrCKwh8bxhgT7d/eI7P4U= github.com/natefinch/pie v0.0.0-20170715172608-9a0d72014007 h1:Ohgj9L0EYOgXxkDp+bczlMBiulwmqYzQpvQNUdtt3oc= @@ -484,6 +506,10 @@ github.com/nfnt/resize v0.0.0-20180221191011-83c6a9932646 h1:zYyBkD/k9seD2A7fsi6 github.com/nfnt/resize v0.0.0-20180221191011-83c6a9932646/go.mod h1:jpp1/29i3P1S/RLdc7JQKbRpFeM1dOBd8T9ki5s+AY8= github.com/nu7hatch/gouuid v0.0.0-20131221200532-179d4d0c4d8d h1:VhgPp6v9qf9Agr/56bj7Y/xa04UccTW04VP0Qed4vnQ= github.com/nu7hatch/gouuid v0.0.0-20131221200532-179d4d0c4d8d/go.mod h1:YUTz3bUH2ZwIWBy3CJBeOBEugqcmXREj14T+iG/4k4U= +github.com/opencontainers/go-digest v1.0.0 h1:apOUWs51W5PlhuyGyz9FCeeBIOUDA/6nW8Oi/yOhh5U= +github.com/opencontainers/go-digest v1.0.0/go.mod h1:0JzlMkj0TRzQZfJkVvzbP0HBR3IKzErnv2BNG4W4MAM= +github.com/opencontainers/image-spec v1.0.2 h1:9yCKha/T5XdGtO0q9Q9a6T5NUCsTn/DrBg0D7ufOcFM= +github.com/opencontainers/image-spec v1.0.2/go.mod h1:BtxoFyWECRxE4U/7sNtV5W15zMzWCbyJoFRP3s7yZA0= github.com/orisano/pixelmatch v0.0.0-20220722002657-fb0b55479cde h1:x0TT0RDC7UhAVbbWWBzr41ElhJx5tXPWkIHA2HWPRuw= github.com/orisano/pixelmatch v0.0.0-20220722002657-fb0b55479cde/go.mod h1:nZgzbfBr3hhjoZnS66nKrHmduYNpc34ny7RK4z5/HM0= github.com/pascaldekloe/goe v0.0.0-20180627143212-57f6aae5913c/go.mod h1:lzWF7FIEvWOWxwDKqyGYQf6ZUaNfKdP144TG7ZOy1lc= @@ -647,7 +673,6 @@ go.yaml.in/yaml/v3 v3.0.3/go.mod h1:tBHosrYAkRZjRAOREWbDnBXUf08JOwYq++0QNwQiWzI= golang.org/x/crypto v0.0.0-20180904163835-0709b304e793/go.mod h1:6SG95UA2DQfeDnfUPMdvaQW0Q7yPrPDi9nlGo2tz2b4= golang.org/x/crypto v0.0.0-20181029021203-45a5f77698d3/go.mod h1:6SG95UA2DQfeDnfUPMdvaQW0Q7yPrPDi9nlGo2tz2b4= golang.org/x/crypto v0.0.0-20190308221718-c2843e01d9a2/go.mod h1:djNgcEr1/C05ACkg1iLfiJU5Ep61QUkGW8qpdssI0+w= -golang.org/x/crypto v0.0.0-20190325154230-a5d413f7728c/go.mod h1:djNgcEr1/C05ACkg1iLfiJU5Ep61QUkGW8qpdssI0+w= golang.org/x/crypto v0.0.0-20190510104115-cbcb75029529/go.mod h1:yigFU9vqHzYiE8UmvKecakEJjdnWj3jj499lnFckfCI= golang.org/x/crypto v0.0.0-20190605123033-f99c8df09eb5/go.mod h1:yigFU9vqHzYiE8UmvKecakEJjdnWj3jj499lnFckfCI= golang.org/x/crypto v0.0.0-20190820162420-60c769a6c586/go.mod h1:yigFU9vqHzYiE8UmvKecakEJjdnWj3jj499lnFckfCI= diff --git a/graphql/schema/schema.graphql b/graphql/schema/schema.graphql index edfdecaac8..1992076f44 100644 --- a/graphql/schema/schema.graphql +++ b/graphql/schema/schema.graphql @@ -271,6 +271,9 @@ type Query { # Get everything with minimal metadata + "Returns the database type." + getDatabaseBackend: SQLDatabaseType! + # Version version: Version! diff --git a/graphql/schema/types/sql.graphql b/graphql/schema/types/sql.graphql index 53615d6f93..1188c63bb0 100644 --- a/graphql/schema/types/sql.graphql +++ b/graphql/schema/types/sql.graphql @@ -18,3 +18,8 @@ type SQLExecResult { """ last_insert_id: Int64 } + +enum SQLDatabaseType { + SQLITE + POSTGRES +} diff --git a/internal/api/resolver.go b/internal/api/resolver.go index 061d0e1a9b..16b3dc94ac 100644 --- a/internal/api/resolver.go +++ b/internal/api/resolver.go @@ -9,6 +9,7 @@ import ( "github.com/stashapp/stash/internal/build" "github.com/stashapp/stash/internal/manager" + "github.com/stashapp/stash/pkg/database" "github.com/stashapp/stash/pkg/logger" "github.com/stashapp/stash/pkg/models" "github.com/stashapp/stash/pkg/plugin/hook" @@ -358,6 +359,18 @@ func (r *mutationResolver) QuerySQL(ctx context.Context, sql string, args []inte }, nil } +func (r *queryResolver) GetDatabaseBackend(context.Context) (SQLDatabaseType, error) { + db := manager.GetInstance().Database + switch db.DatabaseBackend() { + case database.PostgresBackend: + return SQLDatabaseTypePostgres, nil + case database.SqliteBackend: + return SQLDatabaseTypeSQLIte, nil + } + + return "", errors.New("unknown database type") +} + // Get scene marker tags which show up under the video. func (r *queryResolver) SceneMarkerTags(ctx context.Context, scene_id string) ([]*SceneMarkerTag, error) { sceneID, err := strconv.Atoi(scene_id) diff --git a/internal/api/resolver_mutation_migrate.go b/internal/api/resolver_mutation_migrate.go index 083d307e9f..5cc2b110e0 100644 --- a/internal/api/resolver_mutation_migrate.go +++ b/internal/api/resolver_mutation_migrate.go @@ -30,7 +30,7 @@ func (r *mutationResolver) MigrateBlobs(ctx context.Context, input MigrateBlobsI mgr := manager.GetInstance() t := &task.MigrateBlobsJob{ TxnManager: mgr.Database, - BlobStore: mgr.Database.Blobs, + BlobStore: mgr.Database.Blobs(), Vacuumer: mgr.Database, DeleteOld: utils.IsTrue(input.DeleteOld), } diff --git a/internal/autotag/integration_test.go b/internal/autotag/integration_test.go index fc83df848b..4b438d13a6 100644 --- a/internal/autotag/integration_test.go +++ b/internal/autotag/integration_test.go @@ -11,7 +11,9 @@ import ( "testing" "github.com/stashapp/stash/internal/manager/config" + "github.com/stashapp/stash/pkg/database" "github.com/stashapp/stash/pkg/models" + "github.com/stashapp/stash/pkg/postgres" "github.com/stashapp/stash/pkg/sqlite" "github.com/stashapp/stash/pkg/txn" @@ -33,42 +35,59 @@ var existingStudioID int const expectedMatchTitle = "expected match" -var db *sqlite.Database +var db database.Database var r models.Repository -func testTeardown(databaseFile string) { - err := db.Close() - +func testTeardown(db database.Database) { + err := db.Remove() if err != nil { panic(err) } +} - err = os.Remove(databaseFile) - if err != nil { - panic(err) +func IsPostgresTest() *string { + if val, ok := os.LookupEnv("PGSQL_TEST"); ok { + return &val } + return nil } -func runTests(m *testing.M) int { - // create the database file - f, err := os.CreateTemp("", "*.sqlite") - if err != nil { - panic(fmt.Sprintf("Could not create temporary file: %s", err.Error())) - } +func getNewDB() { + if val := IsPostgresTest(); val != nil { + fmt.Printf("Postgres backend for tests detected\n") + db = postgres.NewDatabase() - f.Close() - databaseFile := f.Name() - db = sqlite.NewDatabase() - if err := db.Open(databaseFile); err != nil { - panic(fmt.Sprintf("Could not initialize database: %s", err.Error())) + if err := db.Open(*val); err != nil { + panic(fmt.Sprintf("Could not initialize database: %s", err.Error())) + } + } else { + fmt.Printf("SQLite backend for tests detected\n") + db = sqlite.NewDatabase() + + // create the database file + f, err := os.CreateTemp("", "*.sqlite") + if err != nil { + panic(fmt.Sprintf("Could not create temporary file: %s", err.Error())) + } + + f.Close() + databaseFile := f.Name() + + if err := db.Open(databaseFile); err != nil { + panic(fmt.Sprintf("Could not initialize database: %s", err.Error())) + } } +} + +func runTests(m *testing.M) int { + getNewDB() r = db.Repository() // defer close and delete the database - defer testTeardown(databaseFile) + defer testTeardown(db) - err = populateDB() + err := populateDB() if err != nil { panic(fmt.Sprintf("Could not populate database: %s", err.Error())) } else { diff --git a/internal/manager/init.go b/internal/manager/init.go index b4af5eab78..e3f668762b 100644 --- a/internal/manager/init.go +++ b/internal/manager/init.go @@ -14,6 +14,7 @@ import ( "github.com/stashapp/stash/internal/dlna" "github.com/stashapp/stash/internal/log" "github.com/stashapp/stash/internal/manager/config" + "github.com/stashapp/stash/pkg/database" "github.com/stashapp/stash/pkg/ffmpeg" "github.com/stashapp/stash/pkg/fsutil" "github.com/stashapp/stash/pkg/gallery" @@ -23,6 +24,7 @@ import ( "github.com/stashapp/stash/pkg/logger" "github.com/stashapp/stash/pkg/models/paths" "github.com/stashapp/stash/pkg/plugin" + "github.com/stashapp/stash/pkg/postgres" "github.com/stashapp/stash/pkg/scene" "github.com/stashapp/stash/pkg/scraper" "github.com/stashapp/stash/pkg/session" @@ -35,7 +37,15 @@ import ( func Initialize(cfg *config.Config, l *log.Logger) (*Manager, error) { ctx := context.TODO() - db := sqlite.NewDatabase() + var db database.Database + + upperUrl := strings.ToUpper(cfg.GetDatabasePath()) + if strings.HasPrefix(upperUrl, string(database.PostgresBackend)+":") { + db = postgres.NewDatabase() + } else { + db = sqlite.NewDatabase() + } + repo := db.Repository() // start with empty paths @@ -47,29 +57,29 @@ func Initialize(cfg *config.Config, l *log.Logger) (*Manager, error) { pluginCache := plugin.NewCache(cfg) sceneService := &scene.Service{ - File: db.File, - Repository: db.Scene, - MarkerRepository: db.SceneMarker, + File: db.File(), + Repository: db.Scene(), + MarkerRepository: db.SceneMarker(), PluginCache: pluginCache, Paths: mgrPaths, Config: cfg, } imageService := &image.Service{ - File: db.File, - Repository: db.Image, + File: db.File(), + Repository: db.Image(), } galleryService := &gallery.Service{ - Repository: db.Gallery, - ImageFinder: db.Image, + Repository: db.Gallery(), + ImageFinder: db.Image(), ImageService: imageService, - File: db.File, - Folder: db.Folder, + File: db.File(), + Folder: db.Folder(), } groupService := &group.Service{ - Repository: db.Group, + Repository: db.Group(), } sceneServer := &SceneServer{ @@ -228,7 +238,7 @@ func (s *Manager) postInit(ctx context.Context) error { } if err := s.Database.Open(s.Config.GetDatabasePath()); err != nil { - var migrationNeededErr *sqlite.MigrationNeededError + var migrationNeededErr *database.MigrationNeededError if errors.As(err, &migrationNeededErr) { logger.Warn(err) } else { diff --git a/internal/manager/manager.go b/internal/manager/manager.go index f4f3fa6360..cfad9f6fd6 100644 --- a/internal/manager/manager.go +++ b/internal/manager/manager.go @@ -16,6 +16,7 @@ import ( "github.com/stashapp/stash/internal/dlna" "github.com/stashapp/stash/internal/log" "github.com/stashapp/stash/internal/manager/config" + "github.com/stashapp/stash/pkg/database" "github.com/stashapp/stash/pkg/ffmpeg" "github.com/stashapp/stash/pkg/fsutil" "github.com/stashapp/stash/pkg/job" @@ -26,7 +27,6 @@ import ( "github.com/stashapp/stash/pkg/plugin" "github.com/stashapp/stash/pkg/scraper" "github.com/stashapp/stash/pkg/session" - "github.com/stashapp/stash/pkg/sqlite" // register custom migrations _ "github.com/stashapp/stash/pkg/sqlite/migrations" @@ -60,7 +60,7 @@ type Manager struct { DLNAService *dlna.Service - Database *sqlite.Database + Database database.Database Repository models.Repository SceneService SceneService @@ -85,7 +85,7 @@ func (s *Manager) SetBlobStoreOptions() { blobsPath := s.Config.GetBlobsPath() extraBlobsPaths := s.Config.GetExtraBlobsPaths() - s.Database.SetBlobStoreOptions(sqlite.BlobStoreOptions{ + s.Database.SetBlobStoreOptions(database.BlobStoreOptions{ UseFilesystem: storageType == config.BlobStorageTypeFilesystem, UseDatabase: storageType == config.BlobStorageTypeDatabase, Path: blobsPath, diff --git a/internal/manager/task/migrate.go b/internal/manager/task/migrate.go index 95798d301e..0ade9f1af5 100644 --- a/internal/manager/task/migrate.go +++ b/internal/manager/task/migrate.go @@ -7,9 +7,9 @@ import ( "os" "path/filepath" + "github.com/stashapp/stash/pkg/database" "github.com/stashapp/stash/pkg/job" "github.com/stashapp/stash/pkg/logger" - "github.com/stashapp/stash/pkg/sqlite" ) type migrateJobConfig interface { @@ -20,7 +20,7 @@ type migrateJobConfig interface { type MigrateJob struct { BackupPath string Config migrateJobConfig - Database *sqlite.Database + Database database.Database } type databaseSchemaInfo struct { @@ -106,9 +106,7 @@ func (s *MigrateJob) Execute(ctx context.Context, progress *job.Progress) error } func (s *MigrateJob) required() (ret databaseSchemaInfo, err error) { - database := s.Database - - m, err := sqlite.NewMigrator(database) + m, err := s.Database.NewMigrator() if err != nil { return } @@ -128,9 +126,7 @@ func (s *MigrateJob) required() (ret databaseSchemaInfo, err error) { } func (s *MigrateJob) runMigrations(ctx context.Context, progress *job.Progress) error { - database := s.Database - - m, err := sqlite.NewMigrator(database) + m, err := s.Database.NewMigrator() if err != nil { return err } diff --git a/pkg/sqlite/anonymise_test.go b/pkg/database/anonymise_test.go similarity index 56% rename from pkg/sqlite/anonymise_test.go rename to pkg/database/anonymise_test.go index 868224eefe..56cc299fc3 100644 --- a/pkg/sqlite/anonymise_test.go +++ b/pkg/database/anonymise_test.go @@ -1,16 +1,24 @@ -//go:build integration -// +build integration +//go:build db_integration +// +build db_integration -package sqlite_test +package database_test import ( "context" "os" "testing" + "github.com/stashapp/stash/pkg/postgres" "github.com/stashapp/stash/pkg/sqlite" ) +func IsPostgresTest() *string { + if val, ok := os.LookupEnv("PGSQL_TEST"); ok { + return &val + } + return nil +} + func TestAnonymiser_Anonymise(t *testing.T) { f, err := os.CreateTemp("", "*.sqlite") if err != nil { @@ -22,7 +30,13 @@ func TestAnonymiser_Anonymise(t *testing.T) { defer os.Remove(f.Name()) // use existing database - anonymiser, err := sqlite.NewAnonymiser(db, f.Name()) + var anonymiser *sqlite.Anonymiser + if val := IsPostgresTest(); val != nil { + anonymiser, err = postgres.NewAnonymiser(db.(*postgres.Database), f.Name()) + } else { + anonymiser, err = sqlite.NewAnonymiser(db.(*sqlite.Database), f.Name()) + } + if err != nil { t.Errorf("Could not create anonymiser: %v", err) return diff --git a/pkg/database/batch.go b/pkg/database/batch.go new file mode 100644 index 0000000000..beaf669725 --- /dev/null +++ b/pkg/database/batch.go @@ -0,0 +1,20 @@ +package database + +const DefaultBatchSize = 1000 + +// BatchExec executes the provided function in batches of the provided size. +func BatchExec[T any](ids []T, batchSize int, fn func(batch []T) error) error { + for i := 0; i < len(ids); i += batchSize { + end := i + batchSize + if end > len(ids) { + end = len(ids) + } + + batch := ids[i:end] + if err := fn(batch); err != nil { + return err + } + } + + return nil +} diff --git a/pkg/database/blob.go b/pkg/database/blob.go new file mode 100644 index 0000000000..eb9355789e --- /dev/null +++ b/pkg/database/blob.go @@ -0,0 +1,27 @@ +package database + +import ( + "context" +) + +type BlobStoreOptions struct { + // UseFilesystem should be true if blob data should be stored in the filesystem + UseFilesystem bool + // UseDatabase should be true if blob data should be stored in the database + UseDatabase bool + // Path is the filesystem path to use for storing blobs + Path string + // SupplementaryPaths are alternative filesystem paths that will be used to find blobs + // No changes will be made to these filesystems + SupplementaryPaths []string +} + +type BlobStore interface { + Count(ctx context.Context) (int, error) + Delete(ctx context.Context, checksum string) error + EntryExists(ctx context.Context, checksum string) (bool, error) + FindBlobs(ctx context.Context, n uint, lastChecksum string) ([]string, error) + MigrateBlob(ctx context.Context, checksum string, deleteOld bool) error + Read(ctx context.Context, checksum string) ([]byte, error) + Write(ctx context.Context, data []byte) (string, error) +} diff --git a/pkg/sqlite/blob_test.go b/pkg/database/blob_test.go similarity index 93% rename from pkg/sqlite/blob_test.go rename to pkg/database/blob_test.go index 10c2b93fe4..3fde202fac 100644 --- a/pkg/sqlite/blob_test.go +++ b/pkg/database/blob_test.go @@ -1,7 +1,7 @@ -//go:build integration -// +build integration +//go:build db_integration +// +build db_integration -package sqlite_test +package database_test import ( "context" diff --git a/pkg/database/custom_fields.go b/pkg/database/custom_fields.go new file mode 100644 index 0000000000..6679b41b9b --- /dev/null +++ b/pkg/database/custom_fields.go @@ -0,0 +1,13 @@ +package database + +import ( + "context" + + "github.com/stashapp/stash/pkg/models" +) + +type customFieldsStore interface { + GetCustomFields(ctx context.Context, id int) (map[string]interface{}, error) + GetCustomFieldsBulk(ctx context.Context, ids []int) ([]models.CustomFieldMap, error) + SetCustomFields(ctx context.Context, id int, values models.CustomFieldsInput) error +} diff --git a/pkg/sqlite/custom_fields_test.go b/pkg/database/custom_fields_test.go similarity index 97% rename from pkg/sqlite/custom_fields_test.go rename to pkg/database/custom_fields_test.go index 8ee154aecf..a6166f46e1 100644 --- a/pkg/sqlite/custom_fields_test.go +++ b/pkg/database/custom_fields_test.go @@ -1,7 +1,7 @@ -//go:build integration -// +build integration +//go:build db_integration +// +build db_integration -package sqlite_test +package database_test import ( "context" @@ -181,7 +181,7 @@ func TestSetCustomFields(t *testing.T) { } // use performer custom fields store - store := db.Performer + store := db.Performer() id := performerIDs[performerIdx] for _, tt := range tests { diff --git a/pkg/database/database.go b/pkg/database/database.go new file mode 100644 index 0000000000..22761ea085 --- /dev/null +++ b/pkg/database/database.go @@ -0,0 +1,88 @@ +package database + +import ( + "context" + "errors" + "fmt" + + "github.com/stashapp/stash/pkg/models" +) + +type DatabaseType string + +const ( + PostgresBackend DatabaseType = "POSTGRESQL" + SqliteBackend DatabaseType = "SQLITE" +) + +type Database interface { + Analyze(ctx context.Context) error + Anonymise(outPath string) error + AnonymousDatabasePath(backupDirectoryPath string) string + AppSchemaVersion() uint + Backup(backupPath string) (err error) + Begin(ctx context.Context, writable bool) (context.Context, error) + Close() error + Commit(ctx context.Context) error + DatabaseBackupPath(backupDirectoryPath string) string + DatabasePath() string + DatabaseBackend() DatabaseType + ExecSQL(ctx context.Context, query string, args []interface{}) (*int64, *int64, error) + IsLocked(err error) bool + Open(dbPath string) error + Optimise(ctx context.Context) error + QuerySQL(ctx context.Context, query string, args []interface{}) ([]string, [][]interface{}, error) + ReInitialise() error + Ready() error + Remove() error + Repository() models.Repository + Reset() error + RestoreFromBackup(backupPath string) error + Rollback(ctx context.Context) error + RunAllMigrations() error + SetBlobStoreOptions(options BlobStoreOptions) + Vacuum(ctx context.Context) error + Version() uint + WithDatabase(ctx context.Context) (context.Context, error) + NewMigrator() (MigrateStore, error) + + Blobs() BlobStore + File() FileStore + Folder() FolderStore + Image() ImageStore + Gallery() GalleryStore + GalleryChapter() GalleryChapterStore + Scene() SceneStore + SceneMarker() SceneMarkerStore + Performer() PerformerStore + SavedFilter() SavedFilterStore + Studio() StudioStore + Tag() TagStore + Group() GroupStore +} + +var ( + // ErrDatabaseNotInitialized indicates that the database is not + // initialized, usually due to an incomplete configuration. + ErrDatabaseNotInitialized = errors.New("database not initialized") +) + +// ErrMigrationNeeded indicates that a database migration is needed +// before the database can be initialized +type MigrationNeededError struct { + CurrentSchemaVersion uint + RequiredSchemaVersion uint +} + +func (e *MigrationNeededError) Error() string { + return fmt.Sprintf("database schema version %d does not match required schema version %d", e.CurrentSchemaVersion, e.RequiredSchemaVersion) +} + +type MismatchedSchemaVersionError struct { + CurrentSchemaVersion uint + RequiredSchemaVersion uint +} + +func (e *MismatchedSchemaVersionError) Error() string { + return fmt.Sprintf("schema version %d is incompatible with required schema version %d", e.CurrentSchemaVersion, e.RequiredSchemaVersion) +} diff --git a/pkg/database/date.go b/pkg/database/date.go new file mode 100644 index 0000000000..74aa7790b9 --- /dev/null +++ b/pkg/database/date.go @@ -0,0 +1,79 @@ +package database + +import ( + "database/sql/driver" + "time" + + "github.com/stashapp/stash/pkg/models" + "gopkg.in/guregu/null.v4" +) + +const DateLayout = "2006-01-02" + +// Date represents a date stored as "YYYY-MM-DD" +type Date struct { + Date time.Time +} + +// Scan implements the Scanner interface. +func (d *Date) Scan(value interface{}) error { + d.Date = value.(time.Time) + return nil +} + +// Value implements the driver Valuer interface. +func (d Date) Value() (driver.Value, error) { + return d.Date.Format(DateLayout), nil +} + +// NullDate represents a nullable date stored as "YYYY-MM-DD" +type NullDate struct { + Date time.Time + Valid bool +} + +// Scan implements the Scanner interface. +func (d *NullDate) Scan(value interface{}) error { + var ok bool + d.Date, ok = value.(time.Time) + if !ok { + d.Date = time.Time{} + d.Valid = false + return nil + } + + d.Valid = true + return nil +} + +// Value implements the driver Valuer interface. +func (d NullDate) Value() (driver.Value, error) { + if !d.Valid { + return nil, nil + } + + return d.Date.Format(DateLayout), nil +} + +func (d *NullDate) DatePtr(precision null.Int) *models.Date { + if d == nil || !d.Valid { + return nil + } + + return &models.Date{Time: d.Date, Precision: models.DatePrecision(precision.Int64)} +} + +func NullDateFromDatePtr(d *models.Date) NullDate { + if d == nil { + return NullDate{Valid: false} + } + return NullDate{Date: d.Time, Valid: true} +} + +func DatePrecisionFromDatePtr(d *models.Date) null.Int { + if d == nil { + // default to day precision + return null.Int{} + } + return null.IntFrom(int64(d.Precision)) +} diff --git a/pkg/database/file.go b/pkg/database/file.go new file mode 100644 index 0000000000..3a6ce31118 --- /dev/null +++ b/pkg/database/file.go @@ -0,0 +1,29 @@ +package database + +import ( + "context" + "io/fs" + + "github.com/stashapp/stash/pkg/models" +) + +type FileStore interface { + CountAllInPaths(ctx context.Context, p []string) (int, error) + CountByFolderID(ctx context.Context, folderID models.FolderID) (int, error) + Create(ctx context.Context, f models.File) error + Destroy(ctx context.Context, id models.FileID) error + DestroyFingerprints(ctx context.Context, fileID models.FileID, types []string) error + Find(ctx context.Context, ids ...models.FileID) ([]models.File, error) + FindAllByPath(ctx context.Context, p string, caseSensitive bool) ([]models.File, error) + FindAllInPaths(ctx context.Context, p []string, limit int, offset int) ([]models.File, error) + FindByFileInfo(ctx context.Context, info fs.FileInfo, size int64) ([]models.File, error) + FindByFingerprint(ctx context.Context, fp models.Fingerprint) ([]models.File, error) + FindByPath(ctx context.Context, p string, caseSensitive bool) (models.File, error) + FindByZipFileID(ctx context.Context, zipFileID models.FileID) ([]models.File, error) + GetCaptions(ctx context.Context, fileID models.FileID) ([]*models.VideoCaption, error) + IsPrimary(ctx context.Context, fileID models.FileID) (bool, error) + ModifyFingerprints(ctx context.Context, fileID models.FileID, fingerprints []models.Fingerprint) error + Query(ctx context.Context, options models.FileQueryOptions) (*models.FileQueryResult, error) + Update(ctx context.Context, f models.File) error + UpdateCaptions(ctx context.Context, fileID models.FileID, captions []*models.VideoCaption) error +} diff --git a/pkg/sqlite/file_filter_test.go b/pkg/database/file_filter_test.go similarity index 95% rename from pkg/sqlite/file_filter_test.go rename to pkg/database/file_filter_test.go index 50eed0129e..cf25c25d32 100644 --- a/pkg/sqlite/file_filter_test.go +++ b/pkg/database/file_filter_test.go @@ -1,7 +1,7 @@ -//go:build integration -// +build integration +//go:build db_integration +// +build db_integration -package sqlite_test +package database_test import ( "context" @@ -88,7 +88,7 @@ func TestFileQuery(t *testing.T) { runWithRollbackTxn(t, tt.name, func(t *testing.T, ctx context.Context) { assert := assert.New(t) - results, err := db.File.Query(ctx, models.FileQueryOptions{ + results, err := db.File().Query(ctx, models.FileQueryOptions{ FileFilter: tt.filter, QueryOptions: models.QueryOptions{ FindFilter: tt.findFilter, diff --git a/pkg/sqlite/file_test.go b/pkg/database/file_test.go similarity index 98% rename from pkg/sqlite/file_test.go rename to pkg/database/file_test.go index 8422390c04..6a7a14b9a0 100644 --- a/pkg/sqlite/file_test.go +++ b/pkg/database/file_test.go @@ -1,7 +1,7 @@ -//go:build integration -// +build integration +//go:build db_integration +// +build db_integration -package sqlite_test +package database_test import ( "context" @@ -192,7 +192,7 @@ func Test_fileFileStore_Create(t *testing.T) { }, } - qb := db.File + qb := db.File() for _, tt := range tests { runWithRollbackTxn(t, tt.name, func(t *testing.T, ctx context.Context) { @@ -419,7 +419,7 @@ func Test_fileStore_Update(t *testing.T) { }, } - qb := db.File + qb := db.File() for _, tt := range tests { runWithRollbackTxn(t, tt.name, func(t *testing.T, ctx context.Context) { assert := assert.New(t) @@ -491,7 +491,7 @@ func Test_fileStore_Find(t *testing.T) { }, } - qb := db.File + qb := db.File() for _, tt := range tests { runWithRollbackTxn(t, tt.name, func(t *testing.T, ctx context.Context) { @@ -546,7 +546,7 @@ func Test_FileStore_FindByPath(t *testing.T) { }, } - qb := db.File + qb := db.File() for _, tt := range tests { runWithRollbackTxn(t, tt.name, func(t *testing.T, ctx context.Context) { @@ -598,7 +598,7 @@ func TestFileStore_FindByFingerprint(t *testing.T) { }, } - qb := db.File + qb := db.File() for _, tt := range tests { runWithRollbackTxn(t, tt.name, func(t *testing.T, ctx context.Context) { @@ -647,7 +647,7 @@ func TestFileStore_IsPrimary(t *testing.T) { }, } - qb := db.File + qb := db.File() for _, tt := range tests { runWithRollbackTxn(t, tt.name, func(t *testing.T, ctx context.Context) { diff --git a/pkg/database/folder.go b/pkg/database/folder.go new file mode 100644 index 0000000000..c99fef7ce9 --- /dev/null +++ b/pkg/database/folder.go @@ -0,0 +1,22 @@ +package database + +import ( + "context" + + "github.com/stashapp/stash/pkg/models" +) + +type FolderStore interface { + CountAllInPaths(ctx context.Context, p []string) (int, error) + Create(ctx context.Context, f *models.Folder) error + Destroy(ctx context.Context, id models.FolderID) error + Find(ctx context.Context, id models.FolderID) (*models.Folder, error) + FindAllInPaths(ctx context.Context, p []string, limit int, offset int) ([]*models.Folder, error) + FindByParentFolderID(ctx context.Context, parentFolderID models.FolderID) ([]*models.Folder, error) + FindByPath(ctx context.Context, p string, caseSensitive bool) (*models.Folder, error) + FindByZipFileID(ctx context.Context, zipFileID models.FileID) ([]*models.Folder, error) + Update(ctx context.Context, updatedObject *models.Folder) error + FindByIDs(ctx context.Context, ids []models.FolderID) ([]*models.Folder, error) + FindMany(ctx context.Context, ids []models.FolderID) ([]*models.Folder, error) + Query(ctx context.Context, options models.FolderQueryOptions) (*models.FolderQueryResult, error) +} diff --git a/pkg/sqlite/folder_filter_test.go b/pkg/database/folder_filter_test.go similarity index 94% rename from pkg/sqlite/folder_filter_test.go rename to pkg/database/folder_filter_test.go index c1c7d7a37e..b5e787c846 100644 --- a/pkg/sqlite/folder_filter_test.go +++ b/pkg/database/folder_filter_test.go @@ -1,7 +1,7 @@ -//go:build integration -// +build integration +//go:build db_integration +// +build db_integration -package sqlite_test +package database_test import ( "context" @@ -66,7 +66,7 @@ func TestFolderQuery(t *testing.T) { runWithRollbackTxn(t, tt.name, func(t *testing.T, ctx context.Context) { assert := assert.New(t) - results, err := db.Folder.Query(ctx, models.FolderQueryOptions{ + results, err := db.Folder().Query(ctx, models.FolderQueryOptions{ FolderFilter: tt.filter, QueryOptions: models.QueryOptions{ FindFilter: tt.findFilter, diff --git a/pkg/sqlite/folder_test.go b/pkg/database/folder_test.go similarity index 97% rename from pkg/sqlite/folder_test.go rename to pkg/database/folder_test.go index 15b2b96b83..7a0a901dc1 100644 --- a/pkg/sqlite/folder_test.go +++ b/pkg/database/folder_test.go @@ -1,7 +1,7 @@ -//go:build integration -// +build integration +//go:build db_integration +// +build db_integration -package sqlite_test +package database_test import ( "context" @@ -65,7 +65,7 @@ func Test_FolderStore_Create(t *testing.T) { }, } - qb := db.Folder + qb := db.Folder() for _, tt := range tests { runWithRollbackTxn(t, tt.name, func(t *testing.T, ctx context.Context) { @@ -165,7 +165,7 @@ func Test_FolderStore_Update(t *testing.T) { }, } - qb := db.Folder + qb := db.Folder() for _, tt := range tests { runWithRollbackTxn(t, tt.name, func(t *testing.T, ctx context.Context) { assert := assert.New(t) @@ -224,7 +224,7 @@ func Test_FolderStore_FindByPath(t *testing.T) { }, } - qb := db.Folder + qb := db.Folder() for _, tt := range tests { runWithRollbackTxn(t, tt.name, func(t *testing.T, ctx context.Context) { diff --git a/pkg/database/gallery.go b/pkg/database/gallery.go new file mode 100644 index 0000000000..7d7059fdeb --- /dev/null +++ b/pkg/database/gallery.go @@ -0,0 +1,44 @@ +package database + +import ( + "context" + + "github.com/stashapp/stash/pkg/models" +) + +type GalleryStore interface { + AddFileID(ctx context.Context, id int, fileID models.FileID) error + AddImages(ctx context.Context, galleryID int, imageIDs ...int) error + All(ctx context.Context) ([]*models.Gallery, error) + Count(ctx context.Context) (int, error) + CountByFileID(ctx context.Context, fileID models.FileID) (int, error) + CountByImageID(ctx context.Context, imageID int) (int, error) + Create(ctx context.Context, newObject *models.Gallery, fileIDs []models.FileID) error + Destroy(ctx context.Context, id int) error + Find(ctx context.Context, id int) (*models.Gallery, error) + FindByChecksum(ctx context.Context, checksum string) ([]*models.Gallery, error) + FindByChecksums(ctx context.Context, checksums []string) ([]*models.Gallery, error) + FindByFileID(ctx context.Context, fileID models.FileID) ([]*models.Gallery, error) + FindByFingerprints(ctx context.Context, fp []models.Fingerprint) ([]*models.Gallery, error) + FindByFolderID(ctx context.Context, folderID models.FolderID) ([]*models.Gallery, error) + FindByImageID(ctx context.Context, imageID int) ([]*models.Gallery, error) + FindByPath(ctx context.Context, p string) ([]*models.Gallery, error) + FindBySceneID(ctx context.Context, sceneID int) ([]*models.Gallery, error) + FindMany(ctx context.Context, ids []int) ([]*models.Gallery, error) + FindUserGalleryByTitle(ctx context.Context, title string) ([]*models.Gallery, error) + GetFiles(ctx context.Context, id int) ([]models.File, error) + GetImageIDs(ctx context.Context, galleryID int) ([]int, error) + GetManyFileIDs(ctx context.Context, ids []int) ([][]models.FileID, error) + GetPerformerIDs(ctx context.Context, id int) ([]int, error) + GetSceneIDs(ctx context.Context, id int) ([]int, error) + GetTagIDs(ctx context.Context, id int) ([]int, error) + GetURLs(ctx context.Context, galleryID int) ([]string, error) + Query(ctx context.Context, galleryFilter *models.GalleryFilterType, findFilter *models.FindFilterType) ([]*models.Gallery, int, error) + QueryCount(ctx context.Context, galleryFilter *models.GalleryFilterType, findFilter *models.FindFilterType) (int, error) + RemoveImages(ctx context.Context, galleryID int, imageIDs ...int) error + ResetCover(ctx context.Context, galleryID int) error + SetCover(ctx context.Context, galleryID int, coverImageID int) error + Update(ctx context.Context, updatedObject *models.Gallery) error + UpdateImages(ctx context.Context, galleryID int, imageIDs []int) error + UpdatePartial(ctx context.Context, id int, partial models.GalleryPartial) (*models.Gallery, error) +} diff --git a/pkg/database/gallery_chapter.go b/pkg/database/gallery_chapter.go new file mode 100644 index 0000000000..1b338f18d6 --- /dev/null +++ b/pkg/database/gallery_chapter.go @@ -0,0 +1,17 @@ +package database + +import ( + "context" + + "github.com/stashapp/stash/pkg/models" +) + +type GalleryChapterStore interface { + Create(ctx context.Context, newObject *models.GalleryChapter) error + Destroy(ctx context.Context, id int) error + Find(ctx context.Context, id int) (*models.GalleryChapter, error) + FindByGalleryID(ctx context.Context, galleryID int) ([]*models.GalleryChapter, error) + FindMany(ctx context.Context, ids []int) ([]*models.GalleryChapter, error) + Update(ctx context.Context, updatedObject *models.GalleryChapter) error + UpdatePartial(ctx context.Context, id int, partial models.GalleryChapterPartial) (*models.GalleryChapter, error) +} diff --git a/pkg/sqlite/gallery_chapter_test.go b/pkg/database/gallery_chapter_test.go similarity index 87% rename from pkg/sqlite/gallery_chapter_test.go rename to pkg/database/gallery_chapter_test.go index 4c71ae6b5a..fc10c17820 100644 --- a/pkg/sqlite/gallery_chapter_test.go +++ b/pkg/database/gallery_chapter_test.go @@ -1,7 +1,7 @@ -//go:build integration -// +build integration +//go:build db_integration +// +build db_integration -package sqlite_test +package database_test import ( "context" @@ -12,7 +12,7 @@ import ( func TestChapterFindByGalleryID(t *testing.T) { withTxn(func(ctx context.Context) error { - mqb := db.GalleryChapter + mqb := db.GalleryChapter() galleryID := galleryIDs[galleryIdxWithChapters] chapters, err := mqb.FindByGalleryID(ctx, galleryID) diff --git a/pkg/sqlite/gallery_test.go b/pkg/database/gallery_test.go similarity index 96% rename from pkg/sqlite/gallery_test.go rename to pkg/database/gallery_test.go index 06d7daf17b..9361c73710 100644 --- a/pkg/sqlite/gallery_test.go +++ b/pkg/database/gallery_test.go @@ -1,11 +1,12 @@ -//go:build integration -// +build integration +//go:build db_integration +// +build db_integration -package sqlite_test +package database_test import ( "context" "math" + "sort" "strconv" "testing" "time" @@ -18,27 +19,27 @@ var invalidID = -1 func loadGalleryRelationships(ctx context.Context, expected models.Gallery, actual *models.Gallery) error { if expected.URLs.Loaded() { - if err := actual.LoadURLs(ctx, db.Gallery); err != nil { + if err := actual.LoadURLs(ctx, db.Gallery()); err != nil { return err } } if expected.SceneIDs.Loaded() { - if err := actual.LoadSceneIDs(ctx, db.Gallery); err != nil { + if err := actual.LoadSceneIDs(ctx, db.Gallery()); err != nil { return err } } if expected.TagIDs.Loaded() { - if err := actual.LoadTagIDs(ctx, db.Gallery); err != nil { + if err := actual.LoadTagIDs(ctx, db.Gallery()); err != nil { return err } } if expected.PerformerIDs.Loaded() { - if err := actual.LoadPerformerIDs(ctx, db.Gallery); err != nil { + if err := actual.LoadPerformerIDs(ctx, db.Gallery()); err != nil { return err } } if expected.Files.Loaded() { - if err := actual.LoadFiles(ctx, db.Gallery); err != nil { + if err := actual.LoadFiles(ctx, db.Gallery()); err != nil { return err } } @@ -54,6 +55,19 @@ func loadGalleryRelationships(ctx context.Context, expected models.Gallery, actu return nil } +func sortGallery(copy *models.Gallery) { + // Ordering is not ensured + copy.SceneIDs.Sort() + copy.PerformerIDs.Sort() + copy.TagIDs.Sort() +} + +func sortByID[T any](list []T, getID func(T) int) { + sort.Slice(list, func(i, j int) bool { + return getID(list[i]) < getID(list[j]) + }) +} + func Test_galleryQueryBuilder_Create(t *testing.T) { var ( title = "title" @@ -148,7 +162,7 @@ func Test_galleryQueryBuilder_Create(t *testing.T) { }, } - qb := db.Gallery + qb := db.Gallery() for _, tt := range tests { runWithRollbackTxn(t, tt.name, func(t *testing.T, ctx context.Context) { @@ -180,6 +194,10 @@ func Test_galleryQueryBuilder_Create(t *testing.T) { return } + // Ordering is not ensured + sortGallery(©) + sortGallery(&s) + assert.Equal(copy, s) // ensure can find the scene @@ -198,6 +216,9 @@ func Test_galleryQueryBuilder_Create(t *testing.T) { return } + sortGallery(©) + sortGallery(found) + assert.Equal(copy, *found) return @@ -353,7 +374,7 @@ func Test_galleryQueryBuilder_Update(t *testing.T) { }, } - qb := db.Gallery + qb := db.Gallery() for _, tt := range tests { runWithRollbackTxn(t, tt.name, func(t *testing.T, ctx context.Context) { assert := assert.New(t) @@ -380,6 +401,10 @@ func Test_galleryQueryBuilder_Update(t *testing.T) { return } + // Ordering is not ensured + sortGallery(©) + sortGallery(s) + assert.Equal(copy, *s) return @@ -510,7 +535,7 @@ func Test_galleryQueryBuilder_UpdatePartial(t *testing.T) { }, } for _, tt := range tests { - qb := db.Gallery + qb := db.Gallery() runWithRollbackTxn(t, tt.name, func(t *testing.T, ctx context.Context) { assert := assert.New(t) @@ -779,7 +804,7 @@ func Test_galleryQueryBuilder_UpdatePartialRelationships(t *testing.T) { } for _, tt := range tests { - qb := db.Gallery + qb := db.Gallery() runWithRollbackTxn(t, tt.name, func(t *testing.T, ctx context.Context) { assert := assert.New(t) @@ -809,6 +834,11 @@ func Test_galleryQueryBuilder_UpdatePartialRelationships(t *testing.T) { return } + // Ordering is not ensured + sortGallery(s) + sortGallery(got) + sortGallery(&tt.want) + // only compare fields that were in the partial if tt.partial.PerformerIDs != nil { assert.ElementsMatch(tt.want.PerformerIDs.List(), got.PerformerIDs.List()) @@ -844,7 +874,7 @@ func Test_galleryQueryBuilder_Destroy(t *testing.T) { }, } - qb := db.Gallery + qb := db.Gallery() for _, tt := range tests { runWithRollbackTxn(t, tt.name, func(t *testing.T, ctx context.Context) { @@ -908,7 +938,7 @@ func Test_galleryQueryBuilder_Find(t *testing.T) { }, } - qb := db.Gallery + qb := db.Gallery() for _, tt := range tests { runWithRollbackTxn(t, tt.name, func(t *testing.T, ctx context.Context) { @@ -971,7 +1001,7 @@ func Test_galleryQueryBuilder_FindMany(t *testing.T) { }, } - qb := db.Gallery + qb := db.Gallery() for _, tt := range tests { runWithRollbackTxn(t, tt.name, func(t *testing.T, ctx context.Context) { @@ -1029,7 +1059,7 @@ func Test_galleryQueryBuilder_FindByChecksum(t *testing.T) { }, } - qb := db.Gallery + qb := db.Gallery() for _, tt := range tests { runWithRollbackTxn(t, tt.name, func(t *testing.T, ctx context.Context) { @@ -1092,7 +1122,7 @@ func Test_galleryQueryBuilder_FindByChecksums(t *testing.T) { }, } - qb := db.Gallery + qb := db.Gallery() for _, tt := range tests { runWithRollbackTxn(t, tt.name, func(t *testing.T, ctx context.Context) { @@ -1108,6 +1138,8 @@ func Test_galleryQueryBuilder_FindByChecksums(t *testing.T) { return } + sortByID(tt.want, func(g *models.Gallery) int { return g.ID }) + sortByID(got, func(g *models.Gallery) int { return g.ID }) assert.Equal(tt.want, got) }) } @@ -1150,7 +1182,7 @@ func Test_galleryQueryBuilder_FindByPath(t *testing.T) { }, } - qb := db.Gallery + qb := db.Gallery() for _, tt := range tests { runWithRollbackTxn(t, tt.name, func(t *testing.T, ctx context.Context) { @@ -1192,7 +1224,7 @@ func Test_galleryQueryBuilder_FindBySceneID(t *testing.T) { }, } - qb := db.Gallery + qb := db.Gallery() for _, tt := range tests { runWithRollbackTxn(t, tt.name, func(t *testing.T, ctx context.Context) { @@ -1208,6 +1240,8 @@ func Test_galleryQueryBuilder_FindBySceneID(t *testing.T) { return } + sortByID(tt.want, func(g *models.Gallery) int { return g.ID }) + sortByID(got, func(g *models.Gallery) int { return g.ID }) assert.Equal(tt.want, got) }) } @@ -1237,7 +1271,7 @@ func Test_galleryQueryBuilder_FindByImageID(t *testing.T) { }, } - qb := db.Gallery + qb := db.Gallery() for _, tt := range tests { runWithRollbackTxn(t, tt.name, func(t *testing.T, ctx context.Context) { @@ -1253,6 +1287,8 @@ func Test_galleryQueryBuilder_FindByImageID(t *testing.T) { return } + sortByID(tt.want, func(g *models.Gallery) int { return g.ID }) + sortByID(got, func(g *models.Gallery) int { return g.ID }) assert.Equal(tt.want, got) }) } @@ -1279,7 +1315,7 @@ func Test_galleryQueryBuilder_CountByImageID(t *testing.T) { }, } - qb := db.Gallery + qb := db.Gallery() for _, tt := range tests { runWithRollbackTxn(t, tt.name, func(t *testing.T, ctx context.Context) { @@ -1325,7 +1361,7 @@ func Test_galleryStore_FindByFileID(t *testing.T) { }, } - qb := db.Gallery + qb := db.Gallery() for _, tt := range tests { runWithRollbackTxn(t, tt.name, func(t *testing.T, ctx context.Context) { @@ -1369,7 +1405,7 @@ func Test_galleryStore_FindByFolderID(t *testing.T) { }, } - qb := db.Gallery + qb := db.Gallery() for _, tt := range tests { runWithRollbackTxn(t, tt.name, func(t *testing.T, ctx context.Context) { @@ -1409,7 +1445,7 @@ func TestGalleryQueryQ(t *testing.T) { } func galleryQueryQ(ctx context.Context, t *testing.T, q string, expectedGalleryIdx int) { - qb := db.Gallery + qb := db.Gallery() filter := models.FindFilterType{ Q: &q, @@ -1484,7 +1520,7 @@ func TestGalleryQueryPath(t *testing.T) { }, } - qb := db.Gallery + qb := db.Gallery() for _, tt := range tests { runWithRollbackTxn(t, tt.name, func(t *testing.T, ctx context.Context) { @@ -1511,7 +1547,7 @@ func verifyGalleriesPath(ctx context.Context, t *testing.T, pathCriterion models Path: &pathCriterion, } - sqb := db.Gallery + sqb := db.Gallery() galleries, _, err := sqb.Query(ctx, &galleryFilter, nil) if err != nil { t.Errorf("Error querying gallery: %s", err.Error()) @@ -1545,7 +1581,7 @@ func TestGalleryQueryPathOr(t *testing.T) { } withTxn(func(ctx context.Context) error { - sqb := db.Gallery + sqb := db.Gallery() galleries := queryGallery(ctx, t, sqb, &galleryFilter, nil) @@ -1581,7 +1617,7 @@ func TestGalleryQueryPathAndRating(t *testing.T) { } withTxn(func(ctx context.Context) error { - sqb := db.Gallery + sqb := db.Gallery() galleries := queryGallery(ctx, t, sqb, &galleryFilter, nil) @@ -1621,7 +1657,7 @@ func TestGalleryQueryPathNotRating(t *testing.T) { } withTxn(func(ctx context.Context) error { - sqb := db.Gallery + sqb := db.Gallery() galleries := queryGallery(ctx, t, sqb, &galleryFilter, nil) @@ -1654,7 +1690,7 @@ func TestGalleryIllegalQuery(t *testing.T) { } withTxn(func(ctx context.Context) error { - sqb := db.Gallery + sqb := db.Gallery() _, _, err := sqb.Query(ctx, galleryFilter, nil) assert.NotNil(err) @@ -1720,7 +1756,7 @@ func TestGalleryQueryURL(t *testing.T) { func verifyGalleryQuery(t *testing.T, filter models.GalleryFilterType, verifyFn func(s *models.Gallery)) { withTxn(func(ctx context.Context) error { t.Helper() - sqb := db.Gallery + sqb := db.Gallery() galleries := queryGallery(ctx, t, sqb, &filter, nil) @@ -1768,7 +1804,7 @@ func TestGalleryQueryRating100(t *testing.T) { func verifyGalleriesRating100(t *testing.T, ratingCriterion models.IntCriterionInput) { withTxn(func(ctx context.Context) error { - sqb := db.Gallery + sqb := db.Gallery() galleryFilter := models.GalleryFilterType{ Rating100: &ratingCriterion, } @@ -1788,7 +1824,7 @@ func verifyGalleriesRating100(t *testing.T, ratingCriterion models.IntCriterionI func TestGalleryQueryIsMissingScene(t *testing.T) { withTxn(func(ctx context.Context) error { - qb := db.Gallery + qb := db.Gallery() isMissing := "scenes" galleryFilter := models.GalleryFilterType{ IsMissing: &isMissing, @@ -1832,7 +1868,7 @@ func queryGallery(ctx context.Context, t *testing.T, sqb models.GalleryReader, g func TestGalleryQueryIsMissingStudio(t *testing.T) { withTxn(func(ctx context.Context) error { - sqb := db.Gallery + sqb := db.Gallery() isMissing := "studio" galleryFilter := models.GalleryFilterType{ IsMissing: &isMissing, @@ -1861,7 +1897,7 @@ func TestGalleryQueryIsMissingStudio(t *testing.T) { func TestGalleryQueryIsMissingPerformers(t *testing.T) { withTxn(func(ctx context.Context) error { - sqb := db.Gallery + sqb := db.Gallery() isMissing := "performers" galleryFilter := models.GalleryFilterType{ IsMissing: &isMissing, @@ -1892,7 +1928,7 @@ func TestGalleryQueryIsMissingPerformers(t *testing.T) { func TestGalleryQueryIsMissingTags(t *testing.T) { withTxn(func(ctx context.Context) error { - sqb := db.Gallery + sqb := db.Gallery() isMissing := "tags" galleryFilter := models.GalleryFilterType{ IsMissing: &isMissing, @@ -1918,7 +1954,7 @@ func TestGalleryQueryIsMissingTags(t *testing.T) { func TestGalleryQueryIsMissingDate(t *testing.T) { withTxn(func(ctx context.Context) error { - sqb := db.Gallery + sqb := db.Gallery() isMissing := "date" galleryFilter := models.GalleryFilterType{ IsMissing: &isMissing, @@ -2051,7 +2087,7 @@ func TestGalleryQueryPerformers(t *testing.T) { runWithRollbackTxn(t, tt.name, func(t *testing.T, ctx context.Context) { assert := assert.New(t) - results, _, err := db.Gallery.Query(ctx, &models.GalleryFilterType{ + results, _, err := db.Gallery().Query(ctx, &models.GalleryFilterType{ Performers: &tt.filter, }, nil) if (err != nil) != tt.wantErr { @@ -2187,7 +2223,7 @@ func TestGalleryQueryTags(t *testing.T) { runWithRollbackTxn(t, tt.name, func(t *testing.T, ctx context.Context) { assert := assert.New(t) - results, _, err := db.Gallery.Query(ctx, &models.GalleryFilterType{ + results, _, err := db.Gallery().Query(ctx, &models.GalleryFilterType{ Tags: &tt.filter, }, nil) if (err != nil) != tt.wantErr { @@ -2280,7 +2316,7 @@ func TestGalleryQueryStudio(t *testing.T) { }, } - qb := db.Gallery + qb := db.Gallery() for _, tt := range tests { runWithRollbackTxn(t, tt.name, func(t *testing.T, ctx context.Context) { @@ -2306,7 +2342,7 @@ func TestGalleryQueryStudio(t *testing.T) { func TestGalleryQueryStudioDepth(t *testing.T) { withTxn(func(ctx context.Context) error { - sqb := db.Gallery + sqb := db.Gallery() depth := 2 studioCriterion := models.HierarchicalMultiCriterionInput{ Value: []string{ @@ -2539,7 +2575,7 @@ func TestGalleryQueryPerformerTags(t *testing.T) { runWithRollbackTxn(t, tt.name, func(t *testing.T, ctx context.Context) { assert := assert.New(t) - results, _, err := db.Gallery.Query(ctx, tt.filter, tt.findFilter) + results, _, err := db.Gallery().Query(ctx, tt.filter, tt.findFilter) if (err != nil) != tt.wantErr { t.Errorf("ImageStore.Query() error = %v, wantErr %v", err, tt.wantErr) return @@ -2581,7 +2617,7 @@ func TestGalleryQueryTagCount(t *testing.T) { func verifyGalleriesTagCount(t *testing.T, tagCountCriterion models.IntCriterionInput) { withTxn(func(ctx context.Context) error { - sqb := db.Gallery + sqb := db.Gallery() galleryFilter := models.GalleryFilterType{ TagCount: &tagCountCriterion, } @@ -2622,7 +2658,7 @@ func TestGalleryQueryPerformerCount(t *testing.T) { func verifyGalleriesPerformerCount(t *testing.T, performerCountCriterion models.IntCriterionInput) { withTxn(func(ctx context.Context) error { - sqb := db.Gallery + sqb := db.Gallery() galleryFilter := models.GalleryFilterType{ PerformerCount: &performerCountCriterion, } @@ -2645,7 +2681,7 @@ func verifyGalleriesPerformerCount(t *testing.T, performerCountCriterion models. func TestGalleryQueryAverageResolution(t *testing.T) { withTxn(func(ctx context.Context) error { - qb := db.Gallery + qb := db.Gallery() resolution := models.ResolutionEnumLow galleryFilter := models.GalleryFilterType{ AverageResolution: &models.ResolutionCriterionInput{ @@ -2683,7 +2719,7 @@ func TestGalleryQueryImageCount(t *testing.T) { func verifyGalleriesImageCount(t *testing.T, imageCountCriterion models.IntCriterionInput) { withTxn(func(ctx context.Context) error { - sqb := db.Gallery + sqb := db.Gallery() galleryFilter := models.GalleryFilterType{ ImageCount: &imageCountCriterion, } @@ -2694,7 +2730,7 @@ func verifyGalleriesImageCount(t *testing.T, imageCountCriterion models.IntCrite for _, gallery := range galleries { pp := 0 - result, err := db.Image.Query(ctx, models.ImageQueryOptions{ + result, err := db.Image().Query(ctx, models.ImageQueryOptions{ QueryOptions: models.QueryOptions{ FindFilter: &models.FindFilterType{ PerPage: &pp, @@ -2749,7 +2785,7 @@ func TestGalleryQuerySorting(t *testing.T) { }, } - qb := db.Gallery + qb := db.Gallery() for _, tt := range tests { runWithRollbackTxn(t, tt.name, func(t *testing.T, ctx context.Context) { @@ -2835,7 +2871,7 @@ func TestGalleryStore_AddImages(t *testing.T) { }, } - qb := db.Gallery + qb := db.Gallery() for _, tt := range tests { runWithRollbackTxn(t, tt.name, func(t *testing.T, ctx context.Context) { @@ -2914,7 +2950,7 @@ func TestGalleryStore_RemoveImages(t *testing.T) { }, } - qb := db.Gallery + qb := db.Gallery() for _, tt := range tests { runWithRollbackTxn(t, tt.name, func(t *testing.T, ctx context.Context) { @@ -2944,7 +2980,7 @@ func TestGalleryStore_RemoveImages(t *testing.T) { func TestGalleryQueryHasChapters(t *testing.T) { withTxn(func(ctx context.Context) error { - sqb := db.Gallery + sqb := db.Gallery() hasChapters := "true" galleryFilter := models.GalleryFilterType{ HasChapters: &hasChapters, @@ -2975,25 +3011,25 @@ func TestGalleryQueryHasChapters(t *testing.T) { func TestGallerySetAndResetCover(t *testing.T) { withTxn(func(ctx context.Context) error { - sqb := db.Gallery + sqb := db.Gallery() imagePath2 := getFilePath(folderIdxWithImageFiles, getImageBasename(imageIdx2WithGallery)) - result, err := db.Image.CoverByGalleryID(ctx, galleryIDs[galleryIdxWithTwoImages]) + result, err := db.Image().CoverByGalleryID(ctx, galleryIDs[galleryIdxWithTwoImages]) assert.Nil(t, err) assert.Nil(t, result) err = sqb.SetCover(ctx, galleryIDs[galleryIdxWithTwoImages], imageIDs[imageIdx2WithGallery]) assert.Nil(t, err) - result, err = db.Image.CoverByGalleryID(ctx, galleryIDs[galleryIdxWithTwoImages]) + result, err = db.Image().CoverByGalleryID(ctx, galleryIDs[galleryIdxWithTwoImages]) assert.Nil(t, err) assert.Equal(t, result.Path, imagePath2) err = sqb.ResetCover(ctx, galleryIDs[galleryIdxWithTwoImages]) assert.Nil(t, err) - result, err = db.Image.CoverByGalleryID(ctx, galleryIDs[galleryIdxWithTwoImages]) + result, err = db.Image().CoverByGalleryID(ctx, galleryIDs[galleryIdxWithTwoImages]) assert.Nil(t, err) assert.Nil(t, result) diff --git a/pkg/database/group.go b/pkg/database/group.go new file mode 100644 index 0000000000..0f1d27f237 --- /dev/null +++ b/pkg/database/group.go @@ -0,0 +1,47 @@ +package database + +import ( + "context" + + "github.com/stashapp/stash/pkg/models" +) + +type GroupStore interface { + All(ctx context.Context) ([]*models.Group, error) + Count(ctx context.Context) (int, error) + CountByPerformerID(ctx context.Context, performerID int) (int, error) + CountByStudioID(ctx context.Context, studioID int) (int, error) + Create(ctx context.Context, newObject *models.Group) error + Destroy(ctx context.Context, id int) error + Find(ctx context.Context, id int) (*models.Group, error) + FindByName(ctx context.Context, name string, nocase bool) (*models.Group, error) + FindByNames(ctx context.Context, names []string, nocase bool) ([]*models.Group, error) + FindByPerformerID(ctx context.Context, performerID int) ([]*models.Group, error) + FindByStudioID(ctx context.Context, studioID int) ([]*models.Group, error) + FindInAncestors(ctx context.Context, ascestorIDs []int, ids []int) ([]int, error) + FindMany(ctx context.Context, ids []int) ([]*models.Group, error) + FindSubGroupIDs(ctx context.Context, containingID int, ids []int) ([]int, error) + GetBackImage(ctx context.Context, groupID int) ([]byte, error) + GetFrontImage(ctx context.Context, groupID int) ([]byte, error) + GetURLs(ctx context.Context, groupID int) ([]string, error) + HasBackImage(ctx context.Context, groupID int) (bool, error) + HasFrontImage(ctx context.Context, groupID int) (bool, error) + Query(ctx context.Context, groupFilter *models.GroupFilterType, findFilter *models.FindFilterType) ([]*models.Group, int, error) + QueryCount(ctx context.Context, groupFilter *models.GroupFilterType, findFilter *models.FindFilterType) (int, error) + Update(ctx context.Context, updatedObject *models.Group) error + UpdateBackImage(ctx context.Context, groupID int, backImage []byte) error + UpdateFrontImage(ctx context.Context, groupID int, frontImage []byte) error + UpdatePartial(ctx context.Context, id int, partial models.GroupPartial) (*models.Group, error) + + blobJoinQueryBuilder + tagRelationshipStore + groupRelationshipStore +} + +type groupRelationshipStore interface { + AddSubGroups(ctx context.Context, groupID int, subGroups []models.GroupIDDescription, insertIndex *int) error + GetContainingGroupDescriptions(ctx context.Context, id int) ([]models.GroupIDDescription, error) + GetSubGroupDescriptions(ctx context.Context, id int) ([]models.GroupIDDescription, error) + RemoveSubGroups(ctx context.Context, groupID int, subGroupIDs []int) error + ReorderSubGroups(ctx context.Context, groupID int, subGroupIDs []int, insertPointID int, insertAfter bool) error +} diff --git a/pkg/sqlite/group_test.go b/pkg/database/group_test.go similarity index 98% rename from pkg/sqlite/group_test.go rename to pkg/database/group_test.go index d4a177e86c..e4dcc4bde6 100644 --- a/pkg/sqlite/group_test.go +++ b/pkg/database/group_test.go @@ -1,7 +1,7 @@ -//go:build integration -// +build integration +//go:build db_integration +// +build db_integration -package sqlite_test +package database_test import ( "context" @@ -21,22 +21,22 @@ import ( func loadGroupRelationships(ctx context.Context, expected models.Group, actual *models.Group) error { if expected.URLs.Loaded() { - if err := actual.LoadURLs(ctx, db.Group); err != nil { + if err := actual.LoadURLs(ctx, db.Group()); err != nil { return err } } if expected.TagIDs.Loaded() { - if err := actual.LoadTagIDs(ctx, db.Group); err != nil { + if err := actual.LoadTagIDs(ctx, db.Group()); err != nil { return err } } if expected.ContainingGroups.Loaded() { - if err := actual.LoadContainingGroupIDs(ctx, db.Group); err != nil { + if err := actual.LoadContainingGroupIDs(ctx, db.Group()); err != nil { return err } } if expected.SubGroups.Loaded() { - if err := actual.LoadSubGroupIDs(ctx, db.Group); err != nil { + if err := actual.LoadSubGroupIDs(ctx, db.Group()); err != nil { return err } } @@ -115,7 +115,7 @@ func Test_GroupStore_Create(t *testing.T) { }, } - qb := db.Group + qb := db.Group() for _, tt := range tests { runWithRollbackTxn(t, tt.name, func(t *testing.T, ctx context.Context) { @@ -277,7 +277,7 @@ func Test_groupQueryBuilder_Update(t *testing.T) { }, } - qb := db.Group + qb := db.Group() for _, tt := range tests { runWithRollbackTxn(t, tt.name, func(t *testing.T, ctx context.Context) { assert := assert.New(t) @@ -527,7 +527,7 @@ func Test_groupQueryBuilder_UpdatePartial(t *testing.T) { }, } for _, tt := range tests { - qb := db.Group + qb := db.Group() runWithRollbackTxn(t, tt.name, func(t *testing.T, ctx context.Context) { assert := assert.New(t) @@ -568,7 +568,7 @@ func Test_groupQueryBuilder_UpdatePartial(t *testing.T) { func TestGroupFindByName(t *testing.T) { withTxn(func(ctx context.Context) error { - mqb := db.Group + mqb := db.Group() name := groupNames[groupIdxWithScene] // find a group by name @@ -601,7 +601,7 @@ func TestGroupFindByNames(t *testing.T) { withTxn(func(ctx context.Context) error { var names []string - mqb := db.Group + mqb := db.Group() names = append(names, groupNames[groupIdxWithScene]) // find groups by names @@ -675,7 +675,7 @@ func TestGroupQuery(t *testing.T) { runWithRollbackTxn(t, tt.name, func(t *testing.T, ctx context.Context) { assert := assert.New(t) - results, _, err := db.Group.Query(ctx, tt.filter, tt.findFilter) + results, _, err := db.Group().Query(ctx, tt.filter, tt.findFilter) if (err != nil) != tt.wantErr { t.Errorf("GroupQueryBuilder.Query() error = %v, wantErr %v", err, tt.wantErr) return @@ -697,7 +697,7 @@ func TestGroupQuery(t *testing.T) { func TestGroupQueryStudio(t *testing.T) { withTxn(func(ctx context.Context) error { - mqb := db.Group + mqb := db.Group() studioCriterion := models.HierarchicalMultiCriterionInput{ Value: []string{ strconv.Itoa(studioIDs[studioIdxWithGroup]), @@ -788,7 +788,7 @@ func TestGroupQueryURL(t *testing.T) { func TestGroupQueryURLExcludes(t *testing.T) { withRollbackTxn(func(ctx context.Context) error { - mqb := db.Group + mqb := db.Group() // create group with two URLs group := models.Group{ @@ -839,7 +839,7 @@ func TestGroupQueryURLExcludes(t *testing.T) { func verifyGroupQuery(t *testing.T, filter models.GroupFilterType, verifyFn func(s *models.Group)) { withTxn(func(ctx context.Context) error { t.Helper() - sqb := db.Group + sqb := db.Group() groups := queryGroups(ctx, t, &filter, nil) @@ -861,7 +861,7 @@ func verifyGroupQuery(t *testing.T, filter models.GroupFilterType, verifyFn func } func queryGroups(ctx context.Context, t *testing.T, groupFilter *models.GroupFilterType, findFilter *models.FindFilterType) []*models.Group { - sqb := db.Group + sqb := db.Group() groups, _, err := sqb.Query(ctx, groupFilter, findFilter) if err != nil { t.Errorf("Error querying group: %s", err.Error()) @@ -946,7 +946,7 @@ func TestGroupQueryTagCount(t *testing.T) { func verifyGroupsTagCount(t *testing.T, tagCountCriterion models.IntCriterionInput) { withTxn(func(ctx context.Context) error { - sqb := db.Group + sqb := db.Group() groupFilter := models.GroupFilterType{ TagCount: &tagCountCriterion, } @@ -1011,12 +1011,12 @@ func TestGroupQuerySortOrderIndex(t *testing.T) { withTxn(func(ctx context.Context) error { // just ensure there are no errors - _, _, err := db.Group.Query(ctx, &groupFilter, &findFilter) + _, _, err := db.Group().Query(ctx, &groupFilter, &findFilter) if err != nil { t.Errorf("Error querying group: %s", err.Error()) } - _, _, err = db.Group.Query(ctx, nil, &findFilter) + _, _, err = db.Group().Query(ctx, nil, &findFilter) if err != nil { t.Errorf("Error querying group: %s", err.Error()) } @@ -1027,7 +1027,7 @@ func TestGroupQuerySortOrderIndex(t *testing.T) { func TestGroupUpdateFrontImage(t *testing.T) { if err := withRollbackTxn(func(ctx context.Context) error { - qb := db.Group + qb := db.Group() // create group to test against const name = "TestGroupUpdateGroupImages" @@ -1047,7 +1047,7 @@ func TestGroupUpdateFrontImage(t *testing.T) { func TestGroupUpdateBackImage(t *testing.T) { if err := withRollbackTxn(func(ctx context.Context) error { - qb := db.Group + qb := db.Group() // create group to test against const name = "TestGroupUpdateGroupImages" @@ -1142,7 +1142,7 @@ func TestGroupQueryContainingGroups(t *testing.T) { }, } - qb := db.Group + qb := db.Group() for _, tt := range tests { valueIDs := indexesToIDs(groupIDs, tt.c.valueIdxs) @@ -1255,7 +1255,7 @@ func TestGroupQuerySubGroups(t *testing.T) { }, } - qb := db.Group + qb := db.Group() for _, tt := range tests { valueIDs := indexesToIDs(groupIDs, tt.c.valueIdxs) @@ -1331,7 +1331,7 @@ func TestGroupQueryContainingGroupCount(t *testing.T) { }, } - qb := db.Group + qb := db.Group() for _, tt := range tests { expectedIDs := indexesToIDs(groupIDs, tt.expectedIdxs) @@ -1402,7 +1402,7 @@ func TestGroupQuerySubGroupCount(t *testing.T) { }, } - qb := db.Group + qb := db.Group() for _, tt := range tests { expectedIDs := indexesToIDs(groupIDs, tt.expectedIdxs) @@ -1460,7 +1460,7 @@ func TestGroupFindInAncestors(t *testing.T) { }, } - qb := db.Group + qb := db.Group() for _, tt := range tests { ancestorIDs := indexesToIDs(groupIDs, tt.ancestorIdxs) @@ -1556,7 +1556,7 @@ func TestGroupReorderSubGroups(t *testing.T) { }, } - qb := db.Group + qb := db.Group() for _, tt := range tests { runWithRollbackTxn(t, tt.name, func(t *testing.T, ctx context.Context) { @@ -1666,7 +1666,7 @@ func TestGroupAddSubGroups(t *testing.T) { }, } - qb := db.Group + qb := db.Group() for _, tt := range tests { runWithRollbackTxn(t, tt.name, func(t *testing.T, ctx context.Context) { @@ -1781,7 +1781,7 @@ func TestGroupRemoveSubGroups(t *testing.T) { }, } - qb := db.Group + qb := db.Group() for _, tt := range tests { runWithRollbackTxn(t, tt.name, func(t *testing.T, ctx context.Context) { @@ -1869,7 +1869,7 @@ func TestGroupFindSubGroupIDs(t *testing.T) { }, } - qb := db.Group + qb := db.Group() for _, tt := range tests { runWithRollbackTxn(t, tt.name, func(t *testing.T, ctx context.Context) { diff --git a/pkg/database/image.go b/pkg/database/image.go new file mode 100644 index 0000000000..573a917931 --- /dev/null +++ b/pkg/database/image.go @@ -0,0 +1,51 @@ +package database + +import ( + "context" + + "github.com/stashapp/stash/pkg/models" +) + +type ImageStore interface { + AddFileID(ctx context.Context, id int, fileID models.FileID) error + All(ctx context.Context) ([]*models.Image, error) + Count(ctx context.Context) (int, error) + CountByFileID(ctx context.Context, fileID models.FileID) (int, error) + CountByGalleryID(ctx context.Context, galleryID int) (int, error) + CoverByGalleryID(ctx context.Context, galleryID int) (*models.Image, error) + Create(ctx context.Context, newObject *models.Image, fileIDs []models.FileID) error + Destroy(ctx context.Context, id int) error + Find(ctx context.Context, id int) (*models.Image, error) + FindByChecksum(ctx context.Context, checksum string) ([]*models.Image, error) + FindByFileID(ctx context.Context, fileID models.FileID) ([]*models.Image, error) + FindByFingerprints(ctx context.Context, fp []models.Fingerprint) ([]*models.Image, error) + FindByFolderID(ctx context.Context, folderID models.FolderID) ([]*models.Image, error) + FindByGalleryID(ctx context.Context, galleryID int) ([]*models.Image, error) + FindByGalleryIDIndex(ctx context.Context, galleryID int, index uint) (*models.Image, error) + FindByZipFileID(ctx context.Context, zipFileID models.FileID) ([]*models.Image, error) + FindMany(ctx context.Context, ids []int) ([]*models.Image, error) + GetFiles(ctx context.Context, id int) ([]models.File, error) + GetGalleryIDs(ctx context.Context, imageID int) ([]int, error) + GetManyFileIDs(ctx context.Context, ids []int) ([][]models.FileID, error) + GetPerformerIDs(ctx context.Context, imageID int) ([]int, error) + GetTagIDs(ctx context.Context, imageID int) ([]int, error) + GetURLs(ctx context.Context, imageID int) ([]string, error) + OCount(ctx context.Context) (int, error) + OCountByPerformerID(ctx context.Context, performerID int) (int, error) + Query(ctx context.Context, options models.ImageQueryOptions) (*models.ImageQueryResult, error) + QueryCount(ctx context.Context, imageFilter *models.ImageFilterType, findFilter *models.FindFilterType) (int, error) + RemoveFileID(ctx context.Context, id int, fileID models.FileID) error + Size(ctx context.Context) (float64, error) + Update(ctx context.Context, updatedObject *models.Image) error + UpdatePartial(ctx context.Context, id int, partial models.ImagePartial) (*models.Image, error) + UpdatePerformers(ctx context.Context, imageID int, performerIDs []int) error + UpdateTags(ctx context.Context, imageID int, tagIDs []int) error + OCountByStudioID(ctx context.Context, studioID int) (int, error) + OCountStore +} + +type OCountStore interface { + DecrementOCounter(ctx context.Context, id int) (int, error) + IncrementOCounter(ctx context.Context, id int) (int, error) + ResetOCounter(ctx context.Context, id int) (int, error) +} diff --git a/pkg/sqlite/image_test.go b/pkg/database/image_test.go similarity index 97% rename from pkg/sqlite/image_test.go rename to pkg/database/image_test.go index aa4ed3b99a..21ea4dd4fd 100644 --- a/pkg/sqlite/image_test.go +++ b/pkg/database/image_test.go @@ -1,7 +1,7 @@ -//go:build integration -// +build integration +//go:build db_integration +// +build db_integration -package sqlite_test +package database_test import ( "context" @@ -16,27 +16,27 @@ import ( func loadImageRelationships(ctx context.Context, expected models.Image, actual *models.Image) error { if expected.URLs.Loaded() { - if err := actual.LoadURLs(ctx, db.Image); err != nil { + if err := actual.LoadURLs(ctx, db.Image()); err != nil { return err } } if expected.GalleryIDs.Loaded() { - if err := actual.LoadGalleryIDs(ctx, db.Image); err != nil { + if err := actual.LoadGalleryIDs(ctx, db.Image()); err != nil { return err } } if expected.TagIDs.Loaded() { - if err := actual.LoadTagIDs(ctx, db.Image); err != nil { + if err := actual.LoadTagIDs(ctx, db.Image()); err != nil { return err } } if expected.PerformerIDs.Loaded() { - if err := actual.LoadPerformerIDs(ctx, db.Image); err != nil { + if err := actual.LoadPerformerIDs(ctx, db.Image()); err != nil { return err } } if expected.Files.Loaded() { - if err := actual.LoadFiles(ctx, db.Image); err != nil { + if err := actual.LoadFiles(ctx, db.Image()); err != nil { return err } } @@ -153,7 +153,7 @@ func Test_imageQueryBuilder_Create(t *testing.T) { }, } - qb := db.Image + qb := db.Image() for _, tt := range tests { runWithRollbackTxn(t, tt.name, func(t *testing.T, ctx context.Context) { @@ -360,7 +360,7 @@ func Test_imageQueryBuilder_Update(t *testing.T) { }, } - qb := db.Image + qb := db.Image() for _, tt := range tests { runWithRollbackTxn(t, tt.name, func(t *testing.T, ctx context.Context) { assert := assert.New(t) @@ -511,7 +511,7 @@ func Test_imageQueryBuilder_UpdatePartial(t *testing.T) { }, } for _, tt := range tests { - qb := db.Image + qb := db.Image() runWithRollbackTxn(t, tt.name, func(t *testing.T, ctx context.Context) { assert := assert.New(t) @@ -785,7 +785,7 @@ func Test_imageQueryBuilder_UpdatePartialRelationships(t *testing.T) { } for _, tt := range tests { - qb := db.Image + qb := db.Image() runWithRollbackTxn(t, tt.name, func(t *testing.T, ctx context.Context) { assert := assert.New(t) @@ -853,7 +853,7 @@ func Test_imageQueryBuilder_IncrementOCounter(t *testing.T) { }, } - qb := db.Image + qb := db.Image() for _, tt := range tests { runWithRollbackTxn(t, tt.name, func(t *testing.T, ctx context.Context) { @@ -896,7 +896,7 @@ func Test_imageQueryBuilder_DecrementOCounter(t *testing.T) { }, } - qb := db.Image + qb := db.Image() for _, tt := range tests { runWithRollbackTxn(t, tt.name, func(t *testing.T, ctx context.Context) { @@ -939,7 +939,7 @@ func Test_imageQueryBuilder_ResetOCounter(t *testing.T) { }, } - qb := db.Image + qb := db.Image() for _, tt := range tests { runWithRollbackTxn(t, tt.name, func(t *testing.T, ctx context.Context) { @@ -973,7 +973,7 @@ func Test_imageQueryBuilder_Destroy(t *testing.T) { }, } - qb := db.Image + qb := db.Image() for _, tt := range tests { runWithRollbackTxn(t, tt.name, func(t *testing.T, ctx context.Context) { @@ -1034,7 +1034,7 @@ func Test_imageQueryBuilder_Find(t *testing.T) { }, } - qb := db.Image + qb := db.Image() for _, tt := range tests { runWithRollbackTxn(t, tt.name, func(t *testing.T, ctx context.Context) { @@ -1097,7 +1097,7 @@ func Test_imageQueryBuilder_FindMany(t *testing.T) { }, } - qb := db.Image + qb := db.Image() for _, tt := range tests { runWithRollbackTxn(t, tt.name, func(t *testing.T, ctx context.Context) { @@ -1156,7 +1156,7 @@ func Test_imageQueryBuilder_FindByChecksum(t *testing.T) { }, } - qb := db.Image + qb := db.Image() for _, tt := range tests { runWithRollbackTxn(t, tt.name, func(t *testing.T, ctx context.Context) { @@ -1234,7 +1234,7 @@ func Test_imageQueryBuilder_FindByFingerprints(t *testing.T) { }, } - qb := db.Image + qb := db.Image() for _, tt := range tests { runWithRollbackTxn(t, tt.name, func(t *testing.T, ctx context.Context) { @@ -1276,7 +1276,7 @@ func Test_imageQueryBuilder_FindByGalleryID(t *testing.T) { }, } - qb := db.Image + qb := db.Image() for _, tt := range tests { runWithRollbackTxn(t, tt.name, func(t *testing.T, ctx context.Context) { @@ -1319,7 +1319,7 @@ func Test_imageQueryBuilder_CountByGalleryID(t *testing.T) { }, } - qb := db.Image + qb := db.Image() for _, tt := range tests { runWithRollbackTxn(t, tt.name, func(t *testing.T, ctx context.Context) { @@ -1365,7 +1365,7 @@ func Test_imageStore_FindByFileID(t *testing.T) { }, } - qb := db.Image + qb := db.Image() for _, tt := range tests { runWithRollbackTxn(t, tt.name, func(t *testing.T, ctx context.Context) { @@ -1420,7 +1420,7 @@ func Test_imageStore_FindByFolderID(t *testing.T) { }, } - qb := db.Image + qb := db.Image() for _, tt := range tests { runWithRollbackTxn(t, tt.name, func(t *testing.T, ctx context.Context) { @@ -1469,7 +1469,7 @@ func Test_imageStore_FindByZipFileID(t *testing.T) { }, } - qb := db.Image + qb := db.Image() for _, tt := range tests { runWithRollbackTxn(t, tt.name, func(t *testing.T, ctx context.Context) { @@ -1503,7 +1503,7 @@ func TestImageQueryQ(t *testing.T) { q := getImageStringValue(imageIdx, titleField) - sqb := db.Image + sqb := db.Image() imageQueryQ(ctx, t, sqb, q, imageIdx) @@ -1558,7 +1558,7 @@ func verifyImageQuery(t *testing.T, filter models.ImageFilterType, verifyFn func t.Helper() withTxn(func(ctx context.Context) error { t.Helper() - sqb := db.Image + sqb := db.Image() images := queryImages(ctx, t, sqb, &filter, nil) @@ -1587,7 +1587,7 @@ func TestImageQueryURL(t *testing.T) { verifyFn := func(ctx context.Context, o *models.Image) { t.Helper() - if err := o.LoadURLs(ctx, db.Image); err != nil { + if err := o.LoadURLs(ctx, db.Image()); err != nil { t.Errorf("Error loading scene URLs: %v", err) } @@ -1639,7 +1639,7 @@ func TestImageQueryPath(t *testing.T) { func verifyImagePath(t *testing.T, pathCriterion models.StringCriterionInput, expected int) { withTxn(func(ctx context.Context) error { - sqb := db.Image + sqb := db.Image() imageFilter := models.ImageFilterType{ Path: &pathCriterion, } @@ -1679,7 +1679,7 @@ func TestImageQueryPathOr(t *testing.T) { } withTxn(func(ctx context.Context) error { - sqb := db.Image + sqb := db.Image() images := queryImages(ctx, t, sqb, &imageFilter, nil) @@ -1715,7 +1715,7 @@ func TestImageQueryPathAndRating(t *testing.T) { } withTxn(func(ctx context.Context) error { - sqb := db.Image + sqb := db.Image() images := queryImages(ctx, t, sqb, &imageFilter, nil) @@ -1755,7 +1755,7 @@ func TestImageQueryPathNotRating(t *testing.T) { } withTxn(func(ctx context.Context) error { - sqb := db.Image + sqb := db.Image() images := queryImages(ctx, t, sqb, &imageFilter, nil) @@ -1788,7 +1788,7 @@ func TestImageIllegalQuery(t *testing.T) { } withTxn(func(ctx context.Context) error { - sqb := db.Image + sqb := db.Image() _, _, err := queryImagesWithCount(ctx, sqb, imageFilter, nil) assert.NotNil(err) @@ -1834,7 +1834,7 @@ func TestImageQueryRating100(t *testing.T) { func verifyImagesRating100(t *testing.T, ratingCriterion models.IntCriterionInput) { withTxn(func(ctx context.Context) error { - sqb := db.Image + sqb := db.Image() imageFilter := models.ImageFilterType{ Rating100: &ratingCriterion, } @@ -1873,7 +1873,7 @@ func TestImageQueryOCounter(t *testing.T) { func verifyImagesOCounter(t *testing.T, oCounterCriterion models.IntCriterionInput) { withTxn(func(ctx context.Context) error { - sqb := db.Image + sqb := db.Image() imageFilter := models.ImageFilterType{ OCounter: &oCounterCriterion, } @@ -1902,7 +1902,7 @@ func TestImageQueryResolution(t *testing.T) { func verifyImagesResolution(t *testing.T, resolution models.ResolutionEnum) { withTxn(func(ctx context.Context) error { - sqb := db.Image + sqb := db.Image() imageFilter := models.ImageFilterType{ Resolution: &models.ResolutionCriterionInput{ Value: resolution, @@ -1916,7 +1916,7 @@ func verifyImagesResolution(t *testing.T, resolution models.ResolutionEnum) { } for _, image := range images { - if err := image.LoadPrimaryFile(ctx, db.File); err != nil { + if err := image.LoadPrimaryFile(ctx, db.File()); err != nil { t.Errorf("Error loading primary file: %s", err.Error()) return nil } @@ -1955,7 +1955,7 @@ func verifyImageResolution(t *testing.T, height int, resolution models.Resolutio func TestImageQueryIsMissingGalleries(t *testing.T) { withTxn(func(ctx context.Context) error { - sqb := db.Image + sqb := db.Image() isMissing := "galleries" imageFilter := models.ImageFilterType{ IsMissing: &isMissing, @@ -1992,7 +1992,7 @@ func TestImageQueryIsMissingGalleries(t *testing.T) { func TestImageQueryIsMissingStudio(t *testing.T) { withTxn(func(ctx context.Context) error { - sqb := db.Image + sqb := db.Image() isMissing := "studio" imageFilter := models.ImageFilterType{ IsMissing: &isMissing, @@ -2027,7 +2027,7 @@ func TestImageQueryIsMissingStudio(t *testing.T) { func TestImageQueryIsMissingPerformers(t *testing.T) { withTxn(func(ctx context.Context) error { - sqb := db.Image + sqb := db.Image() isMissing := "performers" imageFilter := models.ImageFilterType{ IsMissing: &isMissing, @@ -2064,7 +2064,7 @@ func TestImageQueryIsMissingPerformers(t *testing.T) { func TestImageQueryIsMissingTags(t *testing.T) { withTxn(func(ctx context.Context) error { - sqb := db.Image + sqb := db.Image() isMissing := "tags" imageFilter := models.ImageFilterType{ IsMissing: &isMissing, @@ -2096,7 +2096,7 @@ func TestImageQueryIsMissingTags(t *testing.T) { func TestImageQueryIsMissingRating(t *testing.T) { withTxn(func(ctx context.Context) error { - sqb := db.Image + sqb := db.Image() isMissing := "rating" imageFilter := models.ImageFilterType{ IsMissing: &isMissing, @@ -2120,7 +2120,7 @@ func TestImageQueryIsMissingRating(t *testing.T) { func TestImageQueryGallery(t *testing.T) { withTxn(func(ctx context.Context) error { - sqb := db.Image + sqb := db.Image() galleryCriterion := models.MultiCriterionInput{ Value: []string{ strconv.Itoa(galleryIDs[galleryIdxWithImage]), @@ -2289,7 +2289,7 @@ func TestImageQueryPerformers(t *testing.T) { runWithRollbackTxn(t, tt.name, func(t *testing.T, ctx context.Context) { assert := assert.New(t) - results, err := db.Image.Query(ctx, models.ImageQueryOptions{ + results, err := db.Image().Query(ctx, models.ImageQueryOptions{ ImageFilter: &models.ImageFilterType{ Performers: &tt.filter, }, @@ -2425,7 +2425,7 @@ func TestImageQueryTags(t *testing.T) { runWithRollbackTxn(t, tt.name, func(t *testing.T, ctx context.Context) { assert := assert.New(t) - results, err := db.Image.Query(ctx, models.ImageQueryOptions{ + results, err := db.Image().Query(ctx, models.ImageQueryOptions{ ImageFilter: &models.ImageFilterType{ Tags: &tt.filter, }, @@ -2518,7 +2518,7 @@ func TestImageQueryStudio(t *testing.T) { }, } - qb := db.Image + qb := db.Image() for _, tt := range tests { runWithRollbackTxn(t, tt.name, func(t *testing.T, ctx context.Context) { @@ -2544,7 +2544,7 @@ func TestImageQueryStudio(t *testing.T) { func TestImageQueryStudioDepth(t *testing.T) { withTxn(func(ctx context.Context) error { - sqb := db.Image + sqb := db.Image() depth := 2 studioCriterion := models.HierarchicalMultiCriterionInput{ Value: []string{ @@ -2786,7 +2786,7 @@ func TestImageQueryPerformerTags(t *testing.T) { runWithRollbackTxn(t, tt.name, func(t *testing.T, ctx context.Context) { assert := assert.New(t) - results, err := db.Image.Query(ctx, models.ImageQueryOptions{ + results, err := db.Image().Query(ctx, models.ImageQueryOptions{ ImageFilter: tt.filter, QueryOptions: models.QueryOptions{ FindFilter: tt.findFilter, @@ -2831,7 +2831,7 @@ func TestImageQueryTagCount(t *testing.T) { func verifyImagesTagCount(t *testing.T, tagCountCriterion models.IntCriterionInput) { withTxn(func(ctx context.Context) error { - sqb := db.Image + sqb := db.Image() imageFilter := models.ImageFilterType{ TagCount: &tagCountCriterion, } @@ -2872,7 +2872,7 @@ func TestImageQueryPerformerCount(t *testing.T) { func verifyImagesPerformerCount(t *testing.T, performerCountCriterion models.IntCriterionInput) { withTxn(func(ctx context.Context) error { - sqb := db.Image + sqb := db.Image() imageFilter := models.ImageFilterType{ PerformerCount: &performerCountCriterion, } @@ -2930,7 +2930,7 @@ func TestImageQuerySorting(t *testing.T) { }, } - qb := db.Image + qb := db.Image() for _, tt := range tests { runWithRollbackTxn(t, tt.name, func(t *testing.T, ctx context.Context) { @@ -2982,7 +2982,7 @@ func TestImageQueryPagination(t *testing.T) { PerPage: &perPage, } - sqb := db.Image + sqb := db.Image() images, _, err := queryImagesWithCount(ctx, sqb, nil, &findFilter) if err != nil { t.Errorf("Error querying image: %s", err.Error()) diff --git a/pkg/database/migrate.go b/pkg/database/migrate.go new file mode 100644 index 0000000000..c80eac3195 --- /dev/null +++ b/pkg/database/migrate.go @@ -0,0 +1,11 @@ +package database + +import "context" + +type MigrateStore interface { + Close() + CurrentSchemaVersion() uint + PostMigrate(ctx context.Context) error + RequiredSchemaVersion() uint + RunMigration(ctx context.Context, newVersion uint) error +} diff --git a/pkg/database/performer.go b/pkg/database/performer.go new file mode 100644 index 0000000000..33f9149814 --- /dev/null +++ b/pkg/database/performer.go @@ -0,0 +1,39 @@ +package database + +import ( + "context" + + "github.com/stashapp/stash/pkg/models" +) + +type PerformerStore interface { + All(ctx context.Context) ([]*models.Performer, error) + Count(ctx context.Context) (int, error) + CountByTagID(ctx context.Context, tagID int) (int, error) + Create(ctx context.Context, newObject *models.CreatePerformerInput) error + Destroy(ctx context.Context, id int) error + DestroyImage(ctx context.Context, id int, blobCol string) error + Find(ctx context.Context, id int) (*models.Performer, error) + FindByGalleryID(ctx context.Context, galleryID int) ([]*models.Performer, error) + FindByImageID(ctx context.Context, imageID int) ([]*models.Performer, error) + FindByNames(ctx context.Context, names []string, nocase bool) ([]*models.Performer, error) + FindBySceneID(ctx context.Context, sceneID int) ([]*models.Performer, error) + FindByStashID(ctx context.Context, stashID models.StashID) ([]*models.Performer, error) + FindByStashIDStatus(ctx context.Context, hasStashID bool, stashboxEndpoint string) ([]*models.Performer, error) + FindMany(ctx context.Context, ids []int) ([]*models.Performer, error) + GetAliases(ctx context.Context, performerID int) ([]string, error) + GetImage(ctx context.Context, performerID int) ([]byte, error) + GetStashIDs(ctx context.Context, performerID int) ([]models.StashID, error) + GetTagIDs(ctx context.Context, id int) ([]int, error) + GetURLs(ctx context.Context, performerID int) ([]string, error) + HasImage(ctx context.Context, performerID int) (bool, error) + Merge(ctx context.Context, source []int, destination int) error + Query(ctx context.Context, performerFilter *models.PerformerFilterType, findFilter *models.FindFilterType) ([]*models.Performer, int, error) + QueryCount(ctx context.Context, performerFilter *models.PerformerFilterType, findFilter *models.FindFilterType) (int, error) + QueryForAutoTag(ctx context.Context, words []string) ([]*models.Performer, error) + Update(ctx context.Context, updatedObject *models.UpdatePerformerInput) error + UpdateImage(ctx context.Context, performerID int, image []byte) error + UpdatePartial(ctx context.Context, id int, partial models.PerformerPartial) (*models.Performer, error) + // blobJoinQueryBuilder + customFieldsStore +} diff --git a/pkg/sqlite/performer_test.go b/pkg/database/performer_test.go similarity index 97% rename from pkg/sqlite/performer_test.go rename to pkg/database/performer_test.go index 8d53ca0dbb..7edae0c352 100644 --- a/pkg/sqlite/performer_test.go +++ b/pkg/database/performer_test.go @@ -1,7 +1,7 @@ -//go:build integration -// +build integration +//go:build db_integration +// +build db_integration -package sqlite_test +package database_test import ( "context" @@ -24,22 +24,22 @@ var testCustomFields = map[string]interface{}{ func loadPerformerRelationships(ctx context.Context, expected models.Performer, actual *models.Performer) error { if expected.Aliases.Loaded() { - if err := actual.LoadAliases(ctx, db.Performer); err != nil { + if err := actual.LoadAliases(ctx, db.Performer()); err != nil { return err } } if expected.URLs.Loaded() { - if err := actual.LoadURLs(ctx, db.Performer); err != nil { + if err := actual.LoadURLs(ctx, db.Performer()); err != nil { return err } } if expected.TagIDs.Loaded() { - if err := actual.LoadTagIDs(ctx, db.Performer); err != nil { + if err := actual.LoadTagIDs(ctx, db.Performer()); err != nil { return err } } if expected.StashIDs.Loaded() { - if err := actual.LoadStashIDs(ctx, db.Performer); err != nil { + if err := actual.LoadStashIDs(ctx, db.Performer()); err != nil { return err } } @@ -76,8 +76,8 @@ func Test_PerformerStore_Create(t *testing.T) { favorite = true endpoint1 = "endpoint1" endpoint2 = "endpoint2" - stashID1 = "stashid1" - stashID2 = "stashid2" + stashID1 = getUUID("stashid1") + stashID2 = getUUID("stashid2") createdAt = time.Date(2001, 1, 1, 0, 0, 0, 0, time.UTC) updatedAt = time.Date(2001, 1, 1, 0, 0, 0, 0, time.UTC) @@ -150,7 +150,7 @@ func Test_PerformerStore_Create(t *testing.T) { }, } - qb := db.Performer + qb := db.Performer() for _, tt := range tests { runWithRollbackTxn(t, tt.name, func(t *testing.T, ctx context.Context) { @@ -239,8 +239,8 @@ func Test_PerformerStore_Update(t *testing.T) { favorite = true endpoint1 = "endpoint1" endpoint2 = "endpoint2" - stashID1 = "stashid1" - stashID2 = "stashid2" + stashID1 = getUUID("stashid1") + stashID2 = getUUID("stashid2") createdAt = time.Date(2001, 1, 1, 0, 0, 0, 0, time.UTC) updatedAt = time.Date(2001, 1, 1, 0, 0, 0, 0, time.UTC) @@ -360,7 +360,7 @@ func Test_PerformerStore_Update(t *testing.T) { }, } - qb := db.Performer + qb := db.Performer() for _, tt := range tests { runWithRollbackTxn(t, tt.name, func(t *testing.T, ctx context.Context) { assert := assert.New(t) @@ -465,8 +465,8 @@ func Test_PerformerStore_UpdatePartial(t *testing.T) { favorite = true endpoint1 = "endpoint1" endpoint2 = "endpoint2" - stashID1 = "stashid1" - stashID2 = "stashid2" + stashID1 = getUUID("stashid1") + stashID2 = getUUID("stashid2") createdAt = time.Date(2001, 1, 1, 0, 0, 0, 0, time.UTC) updatedAt = time.Date(2001, 1, 1, 0, 0, 0, 0, time.UTC) @@ -606,7 +606,7 @@ func Test_PerformerStore_UpdatePartial(t *testing.T) { }, } for _, tt := range tests { - qb := db.Performer + qb := db.Performer() runWithRollbackTxn(t, tt.name, func(t *testing.T, ctx context.Context) { assert := assert.New(t) @@ -691,7 +691,7 @@ func Test_PerformerStore_UpdatePartialCustomFields(t *testing.T) { }, } for _, tt := range tests { - qb := db.Performer + qb := db.Performer() runWithRollbackTxn(t, tt.name, func(t *testing.T, ctx context.Context) { assert := assert.New(t) @@ -719,7 +719,7 @@ func Test_PerformerStore_UpdatePartialCustomFields(t *testing.T) { func TestPerformerFindBySceneID(t *testing.T) { withTxn(func(ctx context.Context) error { - pqb := db.Performer + pqb := db.Performer() sceneID := sceneIDs[sceneIdxWithPerformer] performers, err := pqb.FindBySceneID(ctx, sceneID) @@ -750,7 +750,7 @@ func TestPerformerFindBySceneID(t *testing.T) { func TestPerformerFindByImageID(t *testing.T) { withTxn(func(ctx context.Context) error { - pqb := db.Performer + pqb := db.Performer() imageID := imageIDs[imageIdxWithPerformer] performers, err := pqb.FindByImageID(ctx, imageID) @@ -781,7 +781,7 @@ func TestPerformerFindByImageID(t *testing.T) { func TestPerformerFindByGalleryID(t *testing.T) { withTxn(func(ctx context.Context) error { - pqb := db.Performer + pqb := db.Performer() galleryID := galleryIDs[galleryIdxWithPerformer] performers, err := pqb.FindByGalleryID(ctx, galleryID) @@ -822,7 +822,7 @@ func TestPerformerFindByNames(t *testing.T) { withTxn(func(ctx context.Context) error { var names []string - pqb := db.Performer + pqb := db.Performer() names = append(names, performerNames[performerIdxWithScene]) // find performers by names @@ -1037,7 +1037,7 @@ func TestPerformerIllegalQuery(t *testing.T) { }, } - sqb := db.Performer + sqb := db.Performer() for _, tt := range tests { runWithRollbackTxn(t, tt.name, func(t *testing.T, ctx context.Context) { @@ -1314,7 +1314,7 @@ func TestPerformerQuery(t *testing.T) { runWithRollbackTxn(t, tt.name, func(t *testing.T, ctx context.Context) { assert := assert.New(t) - performers, _, err := db.Performer.Query(ctx, tt.filter, tt.findFilter) + performers, _, err := db.Performer().Query(ctx, tt.filter, tt.findFilter) if (err != nil) != tt.wantErr { t.Errorf("PerformerStore.Query() error = %v, wantErr %v", err, tt.wantErr) return @@ -1550,7 +1550,7 @@ func TestPerformerQueryCustomFields(t *testing.T) { runWithRollbackTxn(t, tt.name, func(t *testing.T, ctx context.Context) { assert := assert.New(t) - performers, _, err := db.Performer.Query(ctx, tt.filter, nil) + performers, _, err := db.Performer().Query(ctx, tt.filter, nil) if (err != nil) != tt.wantErr { t.Errorf("PerformerStore.Query() error = %v, wantErr %v", err, tt.wantErr) return @@ -1633,7 +1633,7 @@ func TestPerformerQueryPenisLength(t *testing.T) { }, } - performers, _, err := db.Performer.Query(ctx, filter, nil) + performers, _, err := db.Performer().Query(ctx, filter, nil) if err != nil { t.Errorf("PerformerStore.Query() error = %v", err) return @@ -1673,7 +1673,7 @@ func verifyFloat(t *testing.T, value *float64, criterion models.FloatCriterionIn func TestPerformerQueryForAutoTag(t *testing.T) { withTxn(func(ctx context.Context) error { - tqb := db.Performer + tqb := db.Performer() name := performerNames[performerIdx1WithScene] // find a performer by name @@ -1693,7 +1693,7 @@ func TestPerformerQueryForAutoTag(t *testing.T) { func TestPerformerUpdatePerformerImage(t *testing.T) { if err := withRollbackTxn(func(ctx context.Context) error { - qb := db.Performer + qb := db.Performer() // create performer to test against const name = "TestPerformerUpdatePerformerImage" @@ -1732,7 +1732,7 @@ func TestPerformerQueryAge(t *testing.T) { func verifyPerformerAge(t *testing.T, ageCriterion models.IntCriterionInput) { withTxn(func(ctx context.Context) error { - qb := db.Performer + qb := db.Performer() performerFilter := models.PerformerFilterType{ Age: &ageCriterion, } @@ -1787,7 +1787,7 @@ func TestPerformerQueryCareerLength(t *testing.T) { func verifyPerformerCareerLength(t *testing.T, criterion models.StringCriterionInput) { withTxn(func(ctx context.Context) error { - qb := db.Performer + qb := db.Performer() performerFilter := models.PerformerFilterType{ CareerLength: &criterion, } @@ -1857,7 +1857,7 @@ func verifyPerformerQuery(t *testing.T, filter models.PerformerFilterType, verif performers := queryPerformers(ctx, t, &filter, nil) for _, performer := range performers { - if err := performer.LoadURLs(ctx, db.Performer); err != nil { + if err := performer.LoadURLs(ctx, db.Performer()); err != nil { t.Errorf("Error loading url relationships: %v", err) } } @@ -1875,7 +1875,7 @@ func verifyPerformerQuery(t *testing.T, filter models.PerformerFilterType, verif func queryPerformers(ctx context.Context, t *testing.T, performerFilter *models.PerformerFilterType, findFilter *models.FindFilterType) []*models.Performer { t.Helper() - performers, _, err := db.Performer.Query(ctx, performerFilter, findFilter) + performers, _, err := db.Performer().Query(ctx, performerFilter, findFilter) if err != nil { t.Errorf("Error querying performers: %s", err.Error()) } @@ -1957,7 +1957,7 @@ func TestPerformerQueryTagCount(t *testing.T) { func verifyPerformersTagCount(t *testing.T, tagCountCriterion models.IntCriterionInput) { withTxn(func(ctx context.Context) error { - sqb := db.Performer + sqb := db.Performer() performerFilter := models.PerformerFilterType{ TagCount: &tagCountCriterion, } @@ -2006,7 +2006,7 @@ func verifyPerformersSceneCount(t *testing.T, sceneCountCriterion models.IntCrit assert.Greater(t, len(performers), 0) for _, performer := range performers { - ids, err := db.Scene.FindByPerformerID(ctx, performer.ID) + ids, err := db.Scene().FindByPerformerID(ctx, performer.ID) if err != nil { return err } @@ -2048,7 +2048,7 @@ func verifyPerformersImageCount(t *testing.T, imageCountCriterion models.IntCrit for _, performer := range performers { pp := 0 - result, err := db.Image.Query(ctx, models.ImageQueryOptions{ + result, err := db.Image().Query(ctx, models.ImageQueryOptions{ QueryOptions: models.QueryOptions{ FindFilter: &models.FindFilterType{ PerPage: &pp, @@ -2103,7 +2103,7 @@ func verifyPerformersGalleryCount(t *testing.T, galleryCountCriterion models.Int for _, performer := range performers { pp := 0 - _, count, err := db.Gallery.Query(ctx, &models.GalleryFilterType{ + _, count, err := db.Gallery().Query(ctx, &models.GalleryFilterType{ Performers: &models.MultiCriterionInput{ Value: []string{strconv.Itoa(performer.ID)}, Modifier: models.CriterionModifierIncludes, @@ -2201,7 +2201,7 @@ func TestPerformerQueryStudio(t *testing.T) { func TestPerformerStashIDs(t *testing.T) { if err := withRollbackTxn(func(ctx context.Context) error { - qb := db.Performer + qb := db.Performer() // create scene to test against const name = "TestPerformerStashIDs" @@ -2228,7 +2228,7 @@ func testPerformerStashIDs(ctx context.Context, t *testing.T, s *models.Performe assert.Len(t, s.StashIDs.List(), 0) // add stash ids - const stashIDStr = "stashID" + var stashIDStr = getUUID("stashID") const endpoint = "endpoint" stashID := models.StashID{ StashID: stashIDStr, @@ -2236,7 +2236,7 @@ func testPerformerStashIDs(ctx context.Context, t *testing.T, s *models.Performe UpdatedAt: epochTime, } - qb := db.Performer + qb := db.Performer() // update stash ids and ensure was updated var err error @@ -2346,7 +2346,7 @@ func TestPerformerQueryIsMissingImage(t *testing.T) { assert.True(t, len(performers) > 0) for _, performer := range performers { - img, err := db.Performer.GetImage(ctx, performer.ID) + img, err := db.Performer().GetImage(ctx, performer.ID) if err != nil { t.Errorf("error getting performer image: %s", err.Error()) } @@ -2364,7 +2364,7 @@ func TestPerformerQueryIsMissingAlias(t *testing.T) { assert.True(t, len(performers) > 0) for _, performer := range performers { - a, err := db.Performer.GetAliases(ctx, performer.ID) + a, err := db.Performer().GetAliases(ctx, performer.ID) if err != nil { t.Errorf("error getting performer aliases: %s", err.Error()) } @@ -2385,7 +2385,7 @@ func TestPerformerQuerySortScenesCount(t *testing.T) { withTxn(func(ctx context.Context) error { // just ensure it queries without error - performers, _, err := db.Performer.Query(ctx, nil, findFilter) + performers, _, err := db.Performer().Query(ctx, nil, findFilter) if err != nil { t.Errorf("Error querying performers: %s", err.Error()) } @@ -2400,7 +2400,7 @@ func TestPerformerQuerySortScenesCount(t *testing.T) { // sort in ascending order direction = models.SortDirectionEnumAsc - performers, _, err = db.Performer.Query(ctx, nil, findFilter) + performers, _, err = db.Performer().Query(ctx, nil, findFilter) if err != nil { t.Errorf("Error querying performers: %s", err.Error()) } @@ -2416,7 +2416,7 @@ func TestPerformerQuerySortScenesCount(t *testing.T) { func TestPerformerCountByTagID(t *testing.T) { withTxn(func(ctx context.Context) error { - sqb := db.Performer + sqb := db.Performer() count, err := sqb.CountByTagID(ctx, tagIDs[tagIdxWithPerformer]) if err != nil { @@ -2439,7 +2439,7 @@ func TestPerformerCountByTagID(t *testing.T) { func TestPerformerCount(t *testing.T) { withTxn(func(ctx context.Context) error { - sqb := db.Performer + sqb := db.Performer() count, err := sqb.Count(ctx) if err != nil { @@ -2454,7 +2454,7 @@ func TestPerformerCount(t *testing.T) { func TestPerformerAll(t *testing.T) { withTxn(func(ctx context.Context) error { - sqb := db.Performer + sqb := db.Performer() all, err := sqb.All(ctx) if err != nil { @@ -2495,7 +2495,7 @@ func TestPerformerStore_FindByStashID(t *testing.T) { { name: "non-existing", stashID: models.StashID{ - StashID: getPerformerStringValue(performerIdxWithScene, "stashid"), + StashID: getUUID("stashid"), Endpoint: "non-existing", }, expectedIDs: []int{}, @@ -2503,7 +2503,7 @@ func TestPerformerStore_FindByStashID(t *testing.T) { }, } - qb := db.Performer + qb := db.Performer() for _, tt := range tests { runWithRollbackTxn(t, tt.name, func(t *testing.T, ctx context.Context) { @@ -2554,7 +2554,7 @@ func TestPerformerStore_FindByStashIDStatus(t *testing.T) { }, } - qb := db.Performer + qb := db.Performer() for _, tt := range tests { runWithRollbackTxn(t, tt.name, func(t *testing.T, ctx context.Context) { diff --git a/pkg/database/saved_filter.go b/pkg/database/saved_filter.go new file mode 100644 index 0000000000..bf645cec80 --- /dev/null +++ b/pkg/database/saved_filter.go @@ -0,0 +1,17 @@ +package database + +import ( + "context" + + "github.com/stashapp/stash/pkg/models" +) + +type SavedFilterStore interface { + All(ctx context.Context) ([]*models.SavedFilter, error) + Create(ctx context.Context, newObject *models.SavedFilter) error + Destroy(ctx context.Context, id int) error + Find(ctx context.Context, id int) (*models.SavedFilter, error) + FindByMode(ctx context.Context, mode models.FilterMode) ([]*models.SavedFilter, error) + FindMany(ctx context.Context, ids []int, ignoreNotFound bool) ([]*models.SavedFilter, error) + Update(ctx context.Context, updatedObject *models.SavedFilter) error +} diff --git a/pkg/sqlite/saved_filter_test.go b/pkg/database/saved_filter_test.go similarity index 82% rename from pkg/sqlite/saved_filter_test.go rename to pkg/database/saved_filter_test.go index 60592a923d..2890b59921 100644 --- a/pkg/sqlite/saved_filter_test.go +++ b/pkg/database/saved_filter_test.go @@ -1,7 +1,7 @@ -//go:build integration -// +build integration +//go:build db_integration +// +build db_integration -package sqlite_test +package database_test import ( "context" @@ -13,7 +13,7 @@ import ( func TestSavedFilterFind(t *testing.T) { withTxn(func(ctx context.Context) error { - savedFilter, err := db.SavedFilter.Find(ctx, savedFilterIDs[savedFilterIdxImage]) + savedFilter, err := db.SavedFilter().Find(ctx, savedFilterIDs[savedFilterIdxImage]) if err != nil { t.Errorf("Error finding saved filter: %s", err.Error()) @@ -27,7 +27,7 @@ func TestSavedFilterFind(t *testing.T) { func TestSavedFilterFindByMode(t *testing.T) { withTxn(func(ctx context.Context) error { - savedFilters, err := db.SavedFilter.FindByMode(ctx, models.FilterModeScenes) + savedFilters, err := db.SavedFilter().FindByMode(ctx, models.FilterModeScenes) if err != nil { t.Errorf("Error finding saved filters: %s", err.Error()) @@ -72,7 +72,7 @@ func TestSavedFilterDestroy(t *testing.T) { ObjectFilter: objectFilter, UIOptions: uiOptions, } - err := db.SavedFilter.Create(ctx, &newFilter) + err := db.SavedFilter().Create(ctx, &newFilter) if err == nil { id = newFilter.ID @@ -82,12 +82,12 @@ func TestSavedFilterDestroy(t *testing.T) { }) withTxn(func(ctx context.Context) error { - return db.SavedFilter.Destroy(ctx, id) + return db.SavedFilter().Destroy(ctx, id) }) // now try to find it withTxn(func(ctx context.Context) error { - found, err := db.SavedFilter.Find(ctx, id) + found, err := db.SavedFilter().Find(ctx, id) if err == nil { assert.Nil(t, found) } diff --git a/pkg/database/scene.go b/pkg/database/scene.go new file mode 100644 index 0000000000..044dee8ff3 --- /dev/null +++ b/pkg/database/scene.go @@ -0,0 +1,96 @@ +package database + +import ( + "context" + "time" + + "github.com/stashapp/stash/pkg/models" +) + +type SceneStore interface { + AddFileID(ctx context.Context, id int, fileID models.FileID) error + AddGalleryIDs(ctx context.Context, sceneID int, galleryIDs []int) error + All(ctx context.Context) ([]*models.Scene, error) + AssignFiles(ctx context.Context, sceneID int, fileIDs []models.FileID) error + Count(ctx context.Context) (int, error) + CountByFileID(ctx context.Context, fileID models.FileID) (int, error) + CountByPerformerID(ctx context.Context, performerID int) (int, error) + CountByStudioID(ctx context.Context, studioID int) (int, error) + CountMissingChecksum(ctx context.Context) (int, error) + CountMissingOSHash(ctx context.Context) (int, error) + Create(ctx context.Context, newObject *models.Scene, fileIDs []models.FileID) error + Destroy(ctx context.Context, id int) error + Duration(ctx context.Context) (float64, error) + Find(ctx context.Context, id int) (*models.Scene, error) + FindByChecksum(ctx context.Context, checksum string) ([]*models.Scene, error) + FindByFileID(ctx context.Context, fileID models.FileID) ([]*models.Scene, error) + FindByFingerprints(ctx context.Context, fp []models.Fingerprint) ([]*models.Scene, error) + FindByGalleryID(ctx context.Context, galleryID int) ([]*models.Scene, error) + FindByGroupID(ctx context.Context, groupID int) ([]*models.Scene, error) + FindByIDs(ctx context.Context, ids []int) ([]*models.Scene, error) + FindByOSHash(ctx context.Context, oshash string) ([]*models.Scene, error) + FindByPath(ctx context.Context, p string) ([]*models.Scene, error) + FindByPerformerID(ctx context.Context, performerID int) ([]*models.Scene, error) + FindByPrimaryFileID(ctx context.Context, fileID models.FileID) ([]*models.Scene, error) + FindDuplicates(ctx context.Context, distance int, durationDiff float64) ([][]*models.Scene, error) + FindMany(ctx context.Context, ids []int) ([]*models.Scene, error) + GetCover(ctx context.Context, sceneID int) ([]byte, error) + GetFiles(ctx context.Context, id int) ([]*models.VideoFile, error) + GetGalleryIDs(ctx context.Context, id int) ([]int, error) + GetGroups(ctx context.Context, id int) (ret []models.GroupsScenes, err error) + GetManyFileIDs(ctx context.Context, ids []int) ([][]models.FileID, error) + GetPerformerIDs(ctx context.Context, id int) ([]int, error) + GetStashIDs(ctx context.Context, sceneID int) ([]models.StashID, error) + GetTagIDs(ctx context.Context, id int) ([]int, error) + GetURLs(ctx context.Context, sceneID int) ([]string, error) + HasCover(ctx context.Context, sceneID int) (bool, error) + OCountByPerformerID(ctx context.Context, performerID int) (int, error) + PlayDuration(ctx context.Context) (float64, error) + Query(ctx context.Context, options models.SceneQueryOptions) (*models.SceneQueryResult, error) + QueryCount(ctx context.Context, sceneFilter *models.SceneFilterType, findFilter *models.FindFilterType) (int, error) + ResetActivity(ctx context.Context, id int, resetResume bool, resetDuration bool) (bool, error) + SaveActivity(ctx context.Context, id int, resumeTime *float64, playDuration *float64) (bool, error) + Size(ctx context.Context) (float64, error) + Update(ctx context.Context, updatedObject *models.Scene) error + UpdateCover(ctx context.Context, sceneID int, image []byte) error + UpdatePartial(ctx context.Context, id int, partial models.ScenePartial) (*models.Scene, error) + Wall(ctx context.Context, q *string) ([]*models.Scene, error) + OCountByGroupID(ctx context.Context, groupID int) (int, error) + OCountByStudioID(ctx context.Context, studioID int) (int, error) + blobJoinQueryBuilder + oDateManager + viewDateManager +} + +type blobJoinQueryBuilder interface { + DestroyImage(ctx context.Context, id int, blobCol string) error + GetImage(ctx context.Context, id int, blobCol string) ([]byte, error) + HasImage(ctx context.Context, id int, blobCol string) (bool, error) + UpdateImage(ctx context.Context, id int, blobCol string, image []byte) error +} + +type oDateManager interface { + AddO(ctx context.Context, id int, dates []time.Time) ([]time.Time, error) + DeleteO(ctx context.Context, id int, dates []time.Time) ([]time.Time, error) + GetAllOCount(ctx context.Context) (int, error) + GetManyOCount(ctx context.Context, ids []int) ([]int, error) + GetManyODates(ctx context.Context, ids []int) ([][]time.Time, error) + GetOCount(ctx context.Context, id int) (int, error) + GetODates(ctx context.Context, id int) ([]time.Time, error) + GetUniqueOCount(ctx context.Context) (int, error) + ResetO(ctx context.Context, id int) (int, error) +} + +type viewDateManager interface { + AddViews(ctx context.Context, id int, dates []time.Time) ([]time.Time, error) + CountAllViews(ctx context.Context) (int, error) + CountUniqueViews(ctx context.Context) (int, error) + CountViews(ctx context.Context, id int) (int, error) + DeleteAllViews(ctx context.Context, id int) (int, error) + DeleteViews(ctx context.Context, id int, dates []time.Time) ([]time.Time, error) + GetManyLastViewed(ctx context.Context, ids []int) ([]*time.Time, error) + GetManyViewCount(ctx context.Context, ids []int) ([]int, error) + GetManyViewDates(ctx context.Context, ids []int) ([][]time.Time, error) + GetViewDates(ctx context.Context, id int) ([]time.Time, error) + LastView(ctx context.Context, id int) (*time.Time, error) +} diff --git a/pkg/database/scene_marker.go b/pkg/database/scene_marker.go new file mode 100644 index 0000000000..fc6f8c9529 --- /dev/null +++ b/pkg/database/scene_marker.go @@ -0,0 +1,26 @@ +package database + +import ( + "context" + + "github.com/stashapp/stash/pkg/models" +) + +type SceneMarkerStore interface { + All(ctx context.Context) ([]*models.SceneMarker, error) + Count(ctx context.Context) (int, error) + CountByTagID(ctx context.Context, tagID int) (int, error) + Create(ctx context.Context, newObject *models.SceneMarker) error + Destroy(ctx context.Context, id int) error + Find(ctx context.Context, id int) (*models.SceneMarker, error) + FindBySceneID(ctx context.Context, sceneID int) ([]*models.SceneMarker, error) + FindMany(ctx context.Context, ids []int) ([]*models.SceneMarker, error) + GetMarkerStrings(ctx context.Context, q *string, sort *string) ([]*models.MarkerStringsResultType, error) + GetTagIDs(ctx context.Context, id int) ([]int, error) + Query(ctx context.Context, sceneMarkerFilter *models.SceneMarkerFilterType, findFilter *models.FindFilterType) ([]*models.SceneMarker, int, error) + QueryCount(ctx context.Context, sceneMarkerFilter *models.SceneMarkerFilterType, findFilter *models.FindFilterType) (int, error) + Update(ctx context.Context, updatedObject *models.SceneMarker) error + UpdatePartial(ctx context.Context, id int, partial models.SceneMarkerPartial) (*models.SceneMarker, error) + UpdateTags(ctx context.Context, id int, tagIDs []int) error + Wall(ctx context.Context, q *string) ([]*models.SceneMarker, error) +} diff --git a/pkg/sqlite/scene_marker_test.go b/pkg/database/scene_marker_test.go similarity index 94% rename from pkg/sqlite/scene_marker_test.go rename to pkg/database/scene_marker_test.go index 64893b3a67..d68ac1a71c 100644 --- a/pkg/sqlite/scene_marker_test.go +++ b/pkg/database/scene_marker_test.go @@ -1,7 +1,7 @@ -//go:build integration -// +build integration +//go:build db_integration +// +build db_integration -package sqlite_test +package database_test import ( "context" @@ -16,7 +16,7 @@ import ( func TestMarkerFindBySceneID(t *testing.T) { withTxn(func(ctx context.Context) error { - mqb := db.SceneMarker + mqb := db.SceneMarker() sceneID := sceneIDs[sceneIdxWithMarkers] markers, err := mqb.FindBySceneID(ctx, sceneID) @@ -44,7 +44,7 @@ func TestMarkerFindBySceneID(t *testing.T) { func TestMarkerCountByTagID(t *testing.T) { withTxn(func(ctx context.Context) error { - mqb := db.SceneMarker + mqb := db.SceneMarker() markerCount, err := mqb.CountByTagID(ctx, tagIDs[tagIdxWithPrimaryMarkers]) @@ -77,7 +77,7 @@ func TestMarkerCountByTagID(t *testing.T) { func TestMarkerQueryQ(t *testing.T) { withTxn(func(ctx context.Context) error { q := getSceneTitle(sceneIdxWithMarkers) - m, _, err := db.SceneMarker.Query(ctx, nil, &models.FindFilterType{ + m, _, err := db.SceneMarker().Query(ctx, nil, &models.FindFilterType{ Q: &q, }) @@ -98,7 +98,7 @@ func TestMarkerQueryQ(t *testing.T) { func TestMarkerQuerySortBySceneUpdated(t *testing.T) { withTxn(func(ctx context.Context) error { sort := "scenes_updated_at" - _, _, err := db.SceneMarker.Query(ctx, nil, &models.FindFilterType{ + _, _, err := db.SceneMarker().Query(ctx, nil, &models.FindFilterType{ Sort: &sort, }) @@ -153,7 +153,7 @@ func TestMarkerQueryTags(t *testing.T) { withTxn(func(ctx context.Context) error { testTags := func(t *testing.T, m *models.SceneMarker, markerFilter *models.SceneMarkerFilterType) { - tagIDs, err := db.SceneMarker.GetTagIDs(ctx, m.ID) + tagIDs, err := db.SceneMarker().GetTagIDs(ctx, m.ID) if err != nil { t.Errorf("error getting marker tag ids: %v", err) } @@ -255,7 +255,7 @@ func TestMarkerQueryTags(t *testing.T) { for _, tc := range cases { t.Run(tc.name, func(t *testing.T) { - markers := queryMarkers(ctx, t, db.SceneMarker, tc.markerFilter, tc.findFilter) + markers := queryMarkers(ctx, t, db.SceneMarker(), tc.markerFilter, tc.findFilter) assert.Greater(t, len(markers), 0) for _, m := range markers { testTags(t, m, tc.markerFilter) @@ -276,13 +276,13 @@ func TestMarkerQuerySceneTags(t *testing.T) { withTxn(func(ctx context.Context) error { testTags := func(t *testing.T, m *models.SceneMarker, markerFilter *models.SceneMarkerFilterType) { - s, err := db.Scene.Find(ctx, m.SceneID) + s, err := db.Scene().Find(ctx, m.SceneID) if err != nil { t.Errorf("error getting marker tag ids: %v", err) return } - if err := s.LoadTagIDs(ctx, db.Scene); err != nil { + if err := s.LoadTagIDs(ctx, db.Scene()); err != nil { t.Errorf("error getting marker tag ids: %v", err) return } @@ -379,7 +379,7 @@ func TestMarkerQuerySceneTags(t *testing.T) { for _, tc := range cases { t.Run(tc.name, func(t *testing.T) { - markers := queryMarkers(ctx, t, db.SceneMarker, tc.markerFilter, tc.findFilter) + markers := queryMarkers(ctx, t, db.SceneMarker(), tc.markerFilter, tc.findFilter) assert.Greater(t, len(markers), 0) for _, m := range markers { testTags(t, m, tc.markerFilter) @@ -475,7 +475,7 @@ func TestMarkerQueryDuration(t *testing.T) { }, } - qb := db.SceneMarker + qb := db.SceneMarker() for _, tt := range cases { runWithRollbackTxn(t, tt.name, func(t *testing.T, ctx context.Context) { diff --git a/pkg/sqlite/scene_test.go b/pkg/database/scene_test.go similarity index 97% rename from pkg/sqlite/scene_test.go rename to pkg/database/scene_test.go index df6676a0f0..2603d2b9d0 100644 --- a/pkg/sqlite/scene_test.go +++ b/pkg/database/scene_test.go @@ -1,7 +1,7 @@ -//go:build integration -// +build integration +//go:build db_integration +// +build db_integration -package sqlite_test +package database_test import ( "context" @@ -22,38 +22,38 @@ import ( func loadSceneRelationships(ctx context.Context, expected models.Scene, actual *models.Scene) error { if expected.URLs.Loaded() { - if err := actual.LoadURLs(ctx, db.Scene); err != nil { + if err := actual.LoadURLs(ctx, db.Scene()); err != nil { return err } } if expected.GalleryIDs.Loaded() { - if err := actual.LoadGalleryIDs(ctx, db.Scene); err != nil { + if err := actual.LoadGalleryIDs(ctx, db.Scene()); err != nil { return err } } if expected.TagIDs.Loaded() { - if err := actual.LoadTagIDs(ctx, db.Scene); err != nil { + if err := actual.LoadTagIDs(ctx, db.Scene()); err != nil { return err } } if expected.PerformerIDs.Loaded() { - if err := actual.LoadPerformerIDs(ctx, db.Scene); err != nil { + if err := actual.LoadPerformerIDs(ctx, db.Scene()); err != nil { return err } } if expected.Groups.Loaded() { - if err := actual.LoadGroups(ctx, db.Scene); err != nil { + if err := actual.LoadGroups(ctx, db.Scene()); err != nil { return err } } if expected.StashIDs.Loaded() { - if err := actual.LoadStashIDs(ctx, db.Scene); err != nil { + if err := actual.LoadStashIDs(ctx, db.Scene()); err != nil { return err } } if expected.Files.Loaded() { - if err := actual.LoadFiles(ctx, db.Scene); err != nil { + if err := actual.LoadFiles(ctx, db.Scene()); err != nil { return err } } @@ -75,6 +75,13 @@ func loadSceneRelationships(ctx context.Context, expected models.Scene, actual * return nil } +func sortScene(copy *models.Scene) { + // Ordering is not ensured + copy.GalleryIDs.Sort() + copy.TagIDs.Sort() + copy.PerformerIDs.Sort() +} + func Test_sceneQueryBuilder_Create(t *testing.T) { var ( title = "title" @@ -91,8 +98,8 @@ func Test_sceneQueryBuilder_Create(t *testing.T) { sceneIndex2 = 234 endpoint1 = "endpoint1" endpoint2 = "endpoint2" - stashID1 = "stashid1" - stashID2 = "stashid2" + stashID1 = getUUID("stashid1") + stashID2 = getUUID("stashid2") date, _ = models.ParseDate("2003-02-01") @@ -237,7 +244,7 @@ func Test_sceneQueryBuilder_Create(t *testing.T) { }, } - qb := db.Scene + qb := db.Scene() for _, tt := range tests { runWithRollbackTxn(t, tt.name, func(t *testing.T, ctx context.Context) { @@ -271,6 +278,8 @@ func Test_sceneQueryBuilder_Create(t *testing.T) { return } + sortScene(©) + sortScene(&s) assert.Equal(copy, s) // ensure can find the scene @@ -288,6 +297,7 @@ func Test_sceneQueryBuilder_Create(t *testing.T) { t.Errorf("loadSceneRelationships() error = %v", err) return } + sortScene(found) assert.Equal(copy, *found) return @@ -325,8 +335,8 @@ func Test_sceneQueryBuilder_Update(t *testing.T) { sceneIndex2 = 234 endpoint1 = "endpoint1" endpoint2 = "endpoint2" - stashID1 = "stashid1" - stashID2 = "stashid2" + stashID1 = getUUID("stashid1") + stashID2 = getUUID("stashid2") date, _ = models.ParseDate("2003-02-01") ) @@ -472,7 +482,7 @@ func Test_sceneQueryBuilder_Update(t *testing.T) { }, } - qb := db.Scene + qb := db.Scene() for _, tt := range tests { runWithRollbackTxn(t, tt.name, func(t *testing.T, ctx context.Context) { assert := assert.New(t) @@ -498,6 +508,8 @@ func Test_sceneQueryBuilder_Update(t *testing.T) { return } + sortScene(©) + sortScene(s) assert.Equal(copy, *s) }) } @@ -537,8 +549,8 @@ func Test_sceneQueryBuilder_UpdatePartial(t *testing.T) { sceneIndex2 = 234 endpoint1 = "endpoint1" endpoint2 = "endpoint2" - stashID1 = "stashid1" - stashID2 = "stashid2" + stashID1 = getUUID("stashid1") + stashID2 = getUUID("stashid2") date, _ = models.ParseDate("2003-02-01") ) @@ -685,7 +697,7 @@ func Test_sceneQueryBuilder_UpdatePartial(t *testing.T) { }, } for _, tt := range tests { - qb := db.Scene + qb := db.Scene() runWithRollbackTxn(t, tt.name, func(t *testing.T, ctx context.Context) { assert := assert.New(t) @@ -709,6 +721,8 @@ func Test_sceneQueryBuilder_UpdatePartial(t *testing.T) { // ignore file ids clearSceneFileIDs(got) + sortScene(&tt.want) + sortScene(got) assert.Equal(tt.want, *got) s, err := qb.Find(ctx, tt.id) @@ -724,6 +738,7 @@ func Test_sceneQueryBuilder_UpdatePartial(t *testing.T) { // ignore file ids clearSceneFileIDs(s) + sortScene(s) assert.Equal(tt.want, *s) }) } @@ -735,8 +750,8 @@ func Test_sceneQueryBuilder_UpdatePartialRelationships(t *testing.T) { sceneIndex2 = 234 endpoint1 = "endpoint1" endpoint2 = "endpoint2" - stashID1 = "stashid1" - stashID2 = "stashid2" + stashID1 = getUUID("stashid1") + stashID2 = getUUID("stashid2") groupScenes = []models.GroupsScenes{ { @@ -1227,7 +1242,7 @@ func Test_sceneQueryBuilder_UpdatePartialRelationships(t *testing.T) { } for _, tt := range tests { - qb := db.Scene + qb := db.Scene() runWithRollbackTxn(t, tt.name, func(t *testing.T, ctx context.Context) { assert := assert.New(t) @@ -1303,7 +1318,7 @@ func Test_sceneQueryBuilder_AddO(t *testing.T) { }, } - qb := db.Scene + qb := db.Scene() for _, tt := range tests { runWithRollbackTxn(t, tt.name, func(t *testing.T, ctx context.Context) { @@ -1340,7 +1355,7 @@ func Test_sceneQueryBuilder_DeleteO(t *testing.T) { }, } - qb := db.Scene + qb := db.Scene() for _, tt := range tests { runWithRollbackTxn(t, tt.name, func(t *testing.T, ctx context.Context) { @@ -1377,7 +1392,7 @@ func Test_sceneQueryBuilder_ResetO(t *testing.T) { }, } - qb := db.Scene + qb := db.Scene() for _, tt := range tests { runWithRollbackTxn(t, tt.name, func(t *testing.T, ctx context.Context) { @@ -1415,7 +1430,7 @@ func Test_sceneQueryBuilder_Destroy(t *testing.T) { }, } - qb := db.Scene + qb := db.Scene() for _, tt := range tests { runWithRollbackTxn(t, tt.name, func(t *testing.T, ctx context.Context) { @@ -1487,7 +1502,7 @@ func Test_sceneQueryBuilder_Find(t *testing.T) { }, } - qb := db.Scene + qb := db.Scene() for _, tt := range tests { runWithRollbackTxn(t, tt.name, func(t *testing.T, ctx context.Context) { @@ -1558,7 +1573,7 @@ func Test_sceneQueryBuilder_FindMany(t *testing.T) { }, } - qb := db.Scene + qb := db.Scene() for _, tt := range tests { runWithRollbackTxn(t, tt.name, func(t *testing.T, ctx context.Context) { @@ -1628,7 +1643,7 @@ func Test_sceneQueryBuilder_FindByChecksum(t *testing.T) { }, } - qb := db.Scene + qb := db.Scene() for _, tt := range tests { runWithRollbackTxn(t, tt.name, func(t *testing.T, ctx context.Context) { @@ -1698,7 +1713,7 @@ func Test_sceneQueryBuilder_FindByOSHash(t *testing.T) { }, } - qb := db.Scene + qb := db.Scene() for _, tt := range tests { runWithRollbackTxn(t, tt.name, func(t *testing.T, ctx context.Context) { @@ -1769,7 +1784,7 @@ func Test_sceneQueryBuilder_FindByPath(t *testing.T) { }, } - qb := db.Scene + qb := db.Scene() for _, tt := range tests { runWithRollbackTxn(t, tt.name, func(t *testing.T, ctx context.Context) { @@ -1811,7 +1826,7 @@ func Test_sceneQueryBuilder_FindByGalleryID(t *testing.T) { }, } - qb := db.Scene + qb := db.Scene() for _, tt := range tests { runWithRollbackTxn(t, tt.name, func(t *testing.T, ctx context.Context) { @@ -1835,7 +1850,7 @@ func Test_sceneQueryBuilder_FindByGalleryID(t *testing.T) { func TestSceneCountByPerformerID(t *testing.T) { withTxn(func(ctx context.Context) error { - sqb := db.Scene + sqb := db.Scene() count, err := sqb.CountByPerformerID(ctx, performerIDs[performerIdxWithScene]) if err != nil { @@ -1886,7 +1901,7 @@ func Test_sceneStore_FindByFileID(t *testing.T) { }, } - qb := db.Scene + qb := db.Scene() for _, tt := range tests { runWithRollbackTxn(t, tt.name, func(t *testing.T, ctx context.Context) { @@ -1932,7 +1947,7 @@ func Test_sceneStore_CountByFileID(t *testing.T) { }, } - qb := db.Scene + qb := db.Scene() for _, tt := range tests { runWithRollbackTxn(t, tt.name, func(t *testing.T, ctx context.Context) { @@ -1959,7 +1974,7 @@ func Test_sceneStore_CountMissingChecksum(t *testing.T) { }, } - qb := db.Scene + qb := db.Scene() for _, tt := range tests { runWithRollbackTxn(t, tt.name, func(t *testing.T, ctx context.Context) { @@ -1986,7 +2001,7 @@ func Test_sceneStore_CountMissingOshash(t *testing.T) { }, } - qb := db.Scene + qb := db.Scene() for _, tt := range tests { runWithRollbackTxn(t, tt.name, func(t *testing.T, ctx context.Context) { @@ -2004,7 +2019,7 @@ func Test_sceneStore_CountMissingOshash(t *testing.T) { func TestSceneWall(t *testing.T) { withTxn(func(ctx context.Context) error { - sqb := db.Scene + sqb := db.Scene() const sceneIdx = 2 wallQuery := getSceneStringValue(sceneIdx, "Details") @@ -2041,7 +2056,7 @@ func TestSceneQueryQ(t *testing.T) { q := getSceneStringValue(sceneIdx, titleField) withTxn(func(ctx context.Context) error { - sqb := db.Scene + sqb := db.Scene() sceneQueryQ(ctx, t, sqb, q, sceneIdx) @@ -2279,7 +2294,7 @@ func TestSceneQuery(t *testing.T) { runWithRollbackTxn(t, tt.name, func(t *testing.T, ctx context.Context) { assert := assert.New(t) - results, err := db.Scene.Query(ctx, models.SceneQueryOptions{ + results, err := db.Scene().Query(ctx, models.SceneQueryOptions{ SceneFilter: tt.filter, QueryOptions: models.QueryOptions{ FindFilter: tt.findFilter, @@ -2392,7 +2407,7 @@ func TestSceneQueryPath(t *testing.T) { }, } - qb := db.Scene + qb := db.Scene() for _, tt := range tests { runWithRollbackTxn(t, tt.name, func(t *testing.T, ctx context.Context) { @@ -2491,7 +2506,7 @@ func TestSceneQueryPathOr(t *testing.T) { } withTxn(func(ctx context.Context) error { - sqb := db.Scene + sqb := db.Scene() scenes := queryScene(ctx, t, sqb, &sceneFilter, nil) @@ -2526,7 +2541,7 @@ func TestSceneQueryPathAndRating(t *testing.T) { } withTxn(func(ctx context.Context) error { - sqb := db.Scene + sqb := db.Scene() scenes := queryScene(ctx, t, sqb, &sceneFilter, nil) @@ -2565,7 +2580,7 @@ func TestSceneQueryPathNotRating(t *testing.T) { } withTxn(func(ctx context.Context) error { - sqb := db.Scene + sqb := db.Scene() scenes := queryScene(ctx, t, sqb, &sceneFilter, nil) @@ -2598,7 +2613,7 @@ func TestSceneIllegalQuery(t *testing.T) { } withTxn(func(ctx context.Context) error { - sqb := db.Scene + sqb := db.Scene() queryOptions := models.SceneQueryOptions{ SceneFilter: sceneFilter, @@ -2625,7 +2640,7 @@ func verifySceneQuery(t *testing.T, filter models.SceneFilterType, verifyFn func t.Helper() withTxn(func(ctx context.Context) error { t.Helper() - sqb := db.Scene + sqb := db.Scene() scenes := queryScene(ctx, t, sqb, &filter, nil) @@ -2648,7 +2663,7 @@ func verifySceneQuery(t *testing.T, filter models.SceneFilterType, verifyFn func func verifyScenesPath(t *testing.T, pathCriterion models.StringCriterionInput) { withTxn(func(ctx context.Context) error { - sqb := db.Scene + sqb := db.Scene() sceneFilter := models.SceneFilterType{ Path: &pathCriterion, } @@ -2757,7 +2772,7 @@ func TestSceneQueryRating100(t *testing.T) { func verifyScenesRating100(t *testing.T, ratingCriterion models.IntCriterionInput) { withTxn(func(ctx context.Context) error { - sqb := db.Scene + sqb := db.Scene() sceneFilter := models.SceneFilterType{ Rating100: &ratingCriterion, } @@ -2816,7 +2831,7 @@ func TestSceneQueryOCounter(t *testing.T) { func verifyScenesOCounter(t *testing.T, oCounterCriterion models.IntCriterionInput) { withTxn(func(ctx context.Context) error { - sqb := db.Scene + sqb := db.Scene() sceneFilter := models.SceneFilterType{ OCounter: &oCounterCriterion, } @@ -2881,7 +2896,7 @@ func TestSceneQueryDuration(t *testing.T) { func verifyScenesDuration(t *testing.T, durationCriterion models.IntCriterionInput) { withTxn(func(ctx context.Context) error { - sqb := db.Scene + sqb := db.Scene() sceneFilter := models.SceneFilterType{ Duration: &durationCriterion, } @@ -2889,7 +2904,7 @@ func verifyScenesDuration(t *testing.T, durationCriterion models.IntCriterionInp scenes := queryScene(ctx, t, sqb, &sceneFilter, nil) for _, scene := range scenes { - if err := scene.LoadPrimaryFile(ctx, db.File); err != nil { + if err := scene.LoadPrimaryFile(ctx, db.File()); err != nil { t.Errorf("Error querying scene files: %v", err) return nil } @@ -2953,7 +2968,7 @@ func TestSceneQueryResolution(t *testing.T) { func verifyScenesResolution(t *testing.T, resolution models.ResolutionEnum) { withTxn(func(ctx context.Context) error { - sqb := db.Scene + sqb := db.Scene() sceneFilter := models.SceneFilterType{ Resolution: &models.ResolutionCriterionInput{ Value: resolution, @@ -2964,7 +2979,7 @@ func verifyScenesResolution(t *testing.T, resolution models.ResolutionEnum) { scenes := queryScene(ctx, t, sqb, &sceneFilter, nil) for _, scene := range scenes { - if err := scene.LoadPrimaryFile(ctx, db.File); err != nil { + if err := scene.LoadPrimaryFile(ctx, db.File()); err != nil { t.Errorf("Error querying scene files: %v", err) return nil } @@ -3016,7 +3031,7 @@ func TestAllResolutionsHaveResolutionRange(t *testing.T) { func TestSceneQueryResolutionModifiers(t *testing.T) { if err := withRollbackTxn(func(ctx context.Context) error { - qb := db.Scene + qb := db.Scene() sceneNoResolution, _ := createScene(ctx, 0, 0) firstScene540P, _ := createScene(ctx, 960, 540) secondScene540P, _ := createScene(ctx, 1280, 719) @@ -3077,13 +3092,13 @@ func createScene(ctx context.Context, width int, height int) (*models.Scene, err Height: height, } - if err := db.File.Create(ctx, sceneFile); err != nil { + if err := db.File().Create(ctx, sceneFile); err != nil { return nil, err } scene := &models.Scene{} - if err := db.Scene.Create(ctx, scene, []models.FileID{sceneFile.ID}); err != nil { + if err := db.Scene().Create(ctx, scene, []models.FileID{sceneFile.ID}); err != nil { return nil, err } @@ -3092,7 +3107,7 @@ func createScene(ctx context.Context, width int, height int) (*models.Scene, err func TestSceneQueryHasMarkers(t *testing.T) { withTxn(func(ctx context.Context) error { - sqb := db.Scene + sqb := db.Scene() hasMarkers := "true" sceneFilter := models.SceneFilterType{ HasMarkers: &hasMarkers, @@ -3128,7 +3143,7 @@ func TestSceneQueryHasMarkers(t *testing.T) { func TestSceneQueryIsMissingGallery(t *testing.T) { withTxn(func(ctx context.Context) error { - sqb := db.Scene + sqb := db.Scene() isMissing := "galleries" sceneFilter := models.SceneFilterType{ IsMissing: &isMissing, @@ -3157,7 +3172,7 @@ func TestSceneQueryIsMissingGallery(t *testing.T) { func TestSceneQueryIsMissingStudio(t *testing.T) { withTxn(func(ctx context.Context) error { - sqb := db.Scene + sqb := db.Scene() isMissing := "studio" sceneFilter := models.SceneFilterType{ IsMissing: &isMissing, @@ -3186,7 +3201,7 @@ func TestSceneQueryIsMissingStudio(t *testing.T) { func TestSceneQueryIsMissingMovies(t *testing.T) { withTxn(func(ctx context.Context) error { - sqb := db.Scene + sqb := db.Scene() isMissing := "movie" sceneFilter := models.SceneFilterType{ IsMissing: &isMissing, @@ -3215,7 +3230,7 @@ func TestSceneQueryIsMissingMovies(t *testing.T) { func TestSceneQueryIsMissingPerformers(t *testing.T) { withTxn(func(ctx context.Context) error { - sqb := db.Scene + sqb := db.Scene() isMissing := "performers" sceneFilter := models.SceneFilterType{ IsMissing: &isMissing, @@ -3246,7 +3261,7 @@ func TestSceneQueryIsMissingPerformers(t *testing.T) { func TestSceneQueryIsMissingDate(t *testing.T) { withTxn(func(ctx context.Context) error { - sqb := db.Scene + sqb := db.Scene() isMissing := "date" sceneFilter := models.SceneFilterType{ IsMissing: &isMissing, @@ -3268,7 +3283,7 @@ func TestSceneQueryIsMissingDate(t *testing.T) { func TestSceneQueryIsMissingTags(t *testing.T) { withTxn(func(ctx context.Context) error { - sqb := db.Scene + sqb := db.Scene() isMissing := "tags" sceneFilter := models.SceneFilterType{ IsMissing: &isMissing, @@ -3294,7 +3309,7 @@ func TestSceneQueryIsMissingTags(t *testing.T) { func TestSceneQueryIsMissingRating(t *testing.T) { withTxn(func(ctx context.Context) error { - sqb := db.Scene + sqb := db.Scene() isMissing := "rating" sceneFilter := models.SceneFilterType{ IsMissing: &isMissing, @@ -3315,7 +3330,7 @@ func TestSceneQueryIsMissingRating(t *testing.T) { func TestSceneQueryIsMissingPhash(t *testing.T) { withTxn(func(ctx context.Context) error { - sqb := db.Scene + sqb := db.Scene() isMissing := "phash" sceneFilter := models.SceneFilterType{ IsMissing: &isMissing, @@ -3446,7 +3461,7 @@ func TestSceneQueryPerformers(t *testing.T) { runWithRollbackTxn(t, tt.name, func(t *testing.T, ctx context.Context) { assert := assert.New(t) - results, err := db.Scene.Query(ctx, models.SceneQueryOptions{ + results, err := db.Scene().Query(ctx, models.SceneQueryOptions{ SceneFilter: &models.SceneFilterType{ Performers: &tt.filter, }, @@ -3582,7 +3597,7 @@ func TestSceneQueryTags(t *testing.T) { runWithRollbackTxn(t, tt.name, func(t *testing.T, ctx context.Context) { assert := assert.New(t) - results, err := db.Scene.Query(ctx, models.SceneQueryOptions{ + results, err := db.Scene().Query(ctx, models.SceneQueryOptions{ SceneFilter: &models.SceneFilterType{ Tags: &tt.filter, }, @@ -3779,7 +3794,7 @@ func TestSceneQueryPerformerTags(t *testing.T) { runWithRollbackTxn(t, tt.name, func(t *testing.T, ctx context.Context) { assert := assert.New(t) - results, err := db.Scene.Query(ctx, models.SceneQueryOptions{ + results, err := db.Scene().Query(ctx, models.SceneQueryOptions{ SceneFilter: tt.filter, QueryOptions: models.QueryOptions{ FindFilter: tt.findFilter, @@ -3873,7 +3888,7 @@ func TestSceneQueryStudio(t *testing.T) { }, } - qb := db.Scene + qb := db.Scene() for _, tt := range tests { runWithRollbackTxn(t, tt.name, func(t *testing.T, ctx context.Context) { @@ -3899,7 +3914,7 @@ func TestSceneQueryStudio(t *testing.T) { func TestSceneQueryStudioDepth(t *testing.T) { withTxn(func(ctx context.Context) error { - sqb := db.Scene + sqb := db.Scene() depth := 2 studioCriterion := models.HierarchicalMultiCriterionInput{ Value: []string{ @@ -4028,7 +4043,7 @@ func TestSceneGroups(t *testing.T) { findFilter.Q = &tt.q } - results, err := db.Scene.Query(ctx, models.SceneQueryOptions{ + results, err := db.Scene().Query(ctx, models.SceneQueryOptions{ SceneFilter: sceneFilter, QueryOptions: models.QueryOptions{ FindFilter: findFilter, @@ -4053,7 +4068,7 @@ func TestSceneGroups(t *testing.T) { func TestSceneQueryMovies(t *testing.T) { withTxn(func(ctx context.Context) error { - sqb := db.Scene + sqb := db.Scene() movieCriterion := models.MultiCriterionInput{ Value: []string{ strconv.Itoa(groupIDs[groupIdxWithScene]), @@ -4093,7 +4108,7 @@ func TestSceneQueryMovies(t *testing.T) { func TestSceneQueryPhashDuplicated(t *testing.T) { withTxn(func(ctx context.Context) error { - sqb := db.Scene + sqb := db.Scene() duplicated := true phashCriterion := models.PHashDuplicationCriterionInput{ Duplicated: &duplicated, @@ -4211,7 +4226,7 @@ func TestSceneQuerySorting(t *testing.T) { }, } - qb := db.Scene + qb := db.Scene() for _, tt := range tests { runWithRollbackTxn(t, tt.name, func(t *testing.T, ctx context.Context) { @@ -4263,7 +4278,7 @@ func TestSceneQueryPagination(t *testing.T) { } withTxn(func(ctx context.Context) error { - sqb := db.Scene + sqb := db.Scene() scenes := queryScene(ctx, t, sqb, nil, &findFilter) assert.Len(t, scenes, 1) @@ -4311,7 +4326,7 @@ func TestSceneQueryTagCount(t *testing.T) { func verifyScenesTagCount(t *testing.T, tagCountCriterion models.IntCriterionInput) { withTxn(func(ctx context.Context) error { - sqb := db.Scene + sqb := db.Scene() sceneFilter := models.SceneFilterType{ TagCount: &tagCountCriterion, } @@ -4352,7 +4367,7 @@ func TestSceneQueryPerformerCount(t *testing.T) { func verifyScenesPerformerCount(t *testing.T, performerCountCriterion models.IntCriterionInput) { withTxn(func(ctx context.Context) error { - sqb := db.Scene + sqb := db.Scene() sceneFilter := models.SceneFilterType{ PerformerCount: &performerCountCriterion, } @@ -4375,7 +4390,7 @@ func verifyScenesPerformerCount(t *testing.T, performerCountCriterion models.Int func TestFindByMovieID(t *testing.T) { withTxn(func(ctx context.Context) error { - sqb := db.Scene + sqb := db.Scene() scenes, err := sqb.FindByGroupID(ctx, groupIDs[groupIdxWithScene]) @@ -4400,7 +4415,7 @@ func TestFindByMovieID(t *testing.T) { func TestFindByPerformerID(t *testing.T) { withTxn(func(ctx context.Context) error { - sqb := db.Scene + sqb := db.Scene() scenes, err := sqb.FindByPerformerID(ctx, performerIDs[performerIdxWithScene]) @@ -4425,7 +4440,7 @@ func TestFindByPerformerID(t *testing.T) { func TestSceneUpdateSceneCover(t *testing.T) { if err := withTxn(func(ctx context.Context) error { - qb := db.Scene + qb := db.Scene() sceneID := sceneIDs[sceneIdxWithGallery] @@ -4437,7 +4452,7 @@ func TestSceneUpdateSceneCover(t *testing.T) { func TestSceneStashIDs(t *testing.T) { if err := withTxn(func(ctx context.Context) error { - qb := db.Scene + qb := db.Scene() // create scene to test against const name = "TestSceneStashIDs" @@ -4464,7 +4479,7 @@ func testSceneStashIDs(ctx context.Context, t *testing.T, s *models.Scene) { assert.Len(t, s.StashIDs.List(), 0) // add stash ids - const stashIDStr = "stashID" + var stashIDStr = getUUID("stashID") const endpoint = "endpoint" stashID := models.StashID{ StashID: stashIDStr, @@ -4472,7 +4487,7 @@ func testSceneStashIDs(ctx context.Context, t *testing.T, s *models.Scene) { UpdatedAt: epochTime, } - qb := db.Scene + qb := db.Scene() // update stash ids and ensure was updated var err error @@ -4514,7 +4529,7 @@ func testSceneStashIDs(ctx context.Context, t *testing.T, s *models.Scene) { func TestSceneQueryQTrim(t *testing.T) { if err := withTxn(func(ctx context.Context) error { - qb := db.Scene + qb := db.Scene() expectedID := sceneIDs[sceneIdxWithSpacedName] @@ -4556,7 +4571,7 @@ func TestSceneQueryQTrim(t *testing.T) { } func TestSceneStore_All(t *testing.T) { - qb := db.Scene + qb := db.Scene() withRollbackTxn(func(ctx context.Context) error { got, err := qb.All(ctx) @@ -4573,7 +4588,7 @@ func TestSceneStore_All(t *testing.T) { } func TestSceneStore_FindDuplicates(t *testing.T) { - qb := db.Scene + qb := db.Scene() withRollbackTxn(func(ctx context.Context) error { distance := 0 @@ -4627,7 +4642,7 @@ func TestSceneStore_AssignFiles(t *testing.T) { }, } - qb := db.Scene + qb := db.Scene() for _, tt := range tests { t.Run(tt.name, func(t *testing.T) { @@ -4663,7 +4678,7 @@ func TestSceneStore_AddView(t *testing.T) { }, } - qb := db.Scene + qb := db.Scene() for _, tt := range tests { t.Run(tt.name, func(t *testing.T) { @@ -4754,7 +4769,7 @@ func TestSceneStore_SaveActivity(t *testing.T) { }, } - qb := db.Scene + qb := db.Scene() for _, tt := range tests { t.Run(tt.name, func(t *testing.T) { @@ -4806,7 +4821,7 @@ func TestSceneStore_SaveActivity(t *testing.T) { // TODO - this should be in history_test and generalised func TestSceneStore_CountAllViews(t *testing.T) { withRollbackTxn(func(ctx context.Context) error { - qb := db.Scene + qb := db.Scene() sceneID := sceneIDs[sceneIdx1WithPerformer] @@ -4839,7 +4854,7 @@ func TestSceneStore_CountAllViews(t *testing.T) { func TestSceneStore_CountUniqueViews(t *testing.T) { withRollbackTxn(func(ctx context.Context) error { - qb := db.Scene + qb := db.Scene() sceneID := sceneIDs[sceneIdx1WithPerformer] diff --git a/pkg/sqlite/setup_test.go b/pkg/database/setup_test.go similarity index 95% rename from pkg/sqlite/setup_test.go rename to pkg/database/setup_test.go index 7e6f821d17..94f51ab3dc 100644 --- a/pkg/sqlite/setup_test.go +++ b/pkg/database/setup_test.go @@ -1,7 +1,7 @@ -//go:build integration -// +build integration +//go:build db_integration +// +build db_integration -package sqlite_test +package database_test import ( "context" @@ -16,12 +16,11 @@ import ( "time" "github.com/stashapp/stash/internal/manager/config" + "github.com/stashapp/stash/pkg/database" "github.com/stashapp/stash/pkg/models" + "github.com/stashapp/stash/pkg/postgres" "github.com/stashapp/stash/pkg/sqlite" "github.com/stashapp/stash/pkg/txn" - - // necessary to register custom migrations - _ "github.com/stashapp/stash/pkg/sqlite/migrations" ) var epochTime = time.Unix(0, 0).UTC() @@ -613,7 +612,7 @@ func indexFromID(ids []int, id int) int { return -1 } -var db *sqlite.Database +var db database.Database func TestMain(m *testing.M) { // initialise empty config - needed by some migrations @@ -646,42 +645,53 @@ func runWithRollbackTxn(t *testing.T, name string, f func(t *testing.T, ctx cont }) } -func testTeardown(databaseFile string) { - err := db.Close() - +func testTeardown(db database.Database) { + err := db.Remove() if err != nil { panic(err) } +} - err = os.Remove(databaseFile) - if err != nil { - panic(err) +func getNewDB() { + if val := IsPostgresTest(); val != nil { + fmt.Printf("Postgres backend for tests detected\n") + db = postgres.NewDatabase() + + if err := db.Open(*val); err != nil { + panic(fmt.Sprintf("Could not initialize database: %s", err.Error())) + } + } else { + fmt.Printf("SQLite backend for tests detected\n") + db = sqlite.NewDatabase() + + // create the database file + f, err := os.CreateTemp("", "*.sqlite") + if err != nil { + panic(fmt.Sprintf("Could not create temporary file: %s", err.Error())) + } + + f.Close() + databaseFile := f.Name() + + if err := db.Open(databaseFile); err != nil { + panic(fmt.Sprintf("Could not initialize database: %s", err.Error())) + } } } func runTests(m *testing.M) int { // create the database file - f, err := os.CreateTemp("", "*.sqlite") - if err != nil { - panic(fmt.Sprintf("Could not create temporary file: %s", err.Error())) - } + getNewDB() - f.Close() - databaseFile := f.Name() - db = sqlite.NewDatabase() - db.SetBlobStoreOptions(sqlite.BlobStoreOptions{ + db.SetBlobStoreOptions(database.BlobStoreOptions{ UseDatabase: true, // don't use filesystem }) - if err := db.Open(databaseFile); err != nil { - panic(fmt.Sprintf("Could not initialize database: %s", err.Error())) - } - // defer close and delete the database - defer testTeardown(databaseFile) + defer testTeardown(db) - err = populateDB() + err := populateDB() if err != nil { panic(fmt.Sprintf("Could not populate database: %s", err.Error())) } @@ -704,11 +714,11 @@ func populateDB() error { return fmt.Errorf("linking folders to zip files: %w", err) } - if err := createTags(ctx, db.Tag, tagsNameCase, tagsNameNoCase); err != nil { + if err := createTags(ctx, db.Tag(), tagsNameCase, tagsNameNoCase); err != nil { return fmt.Errorf("error creating tags: %s", err.Error()) } - if err := createGroups(ctx, db.Group, groupsNameCase, groupsNameNoCase); err != nil { + if err := createGroups(ctx, db.Group(), groupsNameCase, groupsNameNoCase); err != nil { return fmt.Errorf("error creating groups: %s", err.Error()) } @@ -732,15 +742,15 @@ func populateDB() error { return fmt.Errorf("error creating images: %s", err.Error()) } - if err := addTagImage(ctx, db.Tag, tagIdxWithCoverImage); err != nil { + if err := addTagImage(ctx, db.Tag(), tagIdxWithCoverImage); err != nil { return fmt.Errorf("error adding tag image: %s", err.Error()) } - if err := createSavedFilters(ctx, db.SavedFilter, totalSavedFilters); err != nil { + if err := createSavedFilters(ctx, db.SavedFilter(), totalSavedFilters); err != nil { return fmt.Errorf("error creating saved filters: %s", err.Error()) } - if err := linkGroupStudios(ctx, db.Group); err != nil { + if err := linkGroupStudios(ctx, db.Group()); err != nil { return fmt.Errorf("error linking group studios: %s", err.Error()) } @@ -748,21 +758,21 @@ func populateDB() error { return fmt.Errorf("error linking studios parent: %s", err.Error()) } - if err := linkTagsParent(ctx, db.Tag); err != nil { + if err := linkTagsParent(ctx, db.Tag()); err != nil { return fmt.Errorf("error linking tags parent: %s", err.Error()) } - if err := linkGroupsParent(ctx, db.Group); err != nil { + if err := linkGroupsParent(ctx, db.Group()); err != nil { return fmt.Errorf("error linking tags parent: %s", err.Error()) } for _, ms := range markerSpecs { - if err := createMarker(ctx, db.SceneMarker, ms); err != nil { + if err := createMarker(ctx, db.SceneMarker(), ms); err != nil { return fmt.Errorf("error creating scene marker: %s", err.Error()) } } for _, cs := range chapterSpecs { - if err := createChapter(ctx, db.GalleryChapter, cs); err != nil { + if err := createChapter(ctx, db.GalleryChapter(), cs); err != nil { return fmt.Errorf("error creating gallery chapter: %s", err.Error()) } } @@ -775,6 +785,11 @@ func populateDB() error { return nil } +func getUUID(_ string) string { + // TODO: Encode input string + return "00000000-0000-0000-0000-000000000000" +} + func getFolderPath(index int, parentFolderIdx *int) string { path := getPrefixedStringValue("folder", index, pathField) @@ -809,7 +824,7 @@ func makeFolder(i int) models.Folder { } func createFolders(ctx context.Context) error { - qb := db.Folder + qb := db.Folder() for i := 0; i < totalFolders; i++ { folder := makeFolder(i) @@ -831,14 +846,14 @@ func linkFoldersToZip(ctx context.Context) error { folderID := folderIDs[folderIdx] fileID := fileIDs[fileIdx] - f, err := db.Folder.Find(ctx, folderID) + f, err := db.Folder().Find(ctx, folderID) if err != nil { return fmt.Errorf("Error finding folder [%d] to link to zip file [%d]", folderID, fileID) } f.ZipFileID = &fileID - if err := db.Folder.Update(ctx, f); err != nil { + if err := db.Folder().Update(ctx, f); err != nil { return fmt.Errorf("Error linking folder [%d] to zip file [%d]: %s", folderIdx, fileIdx, err.Error()) } } @@ -933,7 +948,7 @@ func makeFile(i int) models.File { } func createFiles(ctx context.Context) error { - qb := db.File + qb := db.File() for i := 0; i < totalFiles; i++ { file := makeFile(i) @@ -1183,8 +1198,8 @@ func makeScene(i int) *models.Scene { } func createScenes(ctx context.Context, n int) error { - sqb := db.Scene - fqb := db.File + sqb := db.Scene() + fqb := db.File() for i := 0; i < n; i++ { f := makeSceneFile(i) @@ -1272,8 +1287,8 @@ func makeImage(i int) *models.Image { } func createImages(ctx context.Context, n int) error { - qb := db.Image - fqb := db.File + qb := db.Image() + fqb := db.File() for i := 0; i < n; i++ { f := makeImageFile(i) @@ -1369,8 +1384,8 @@ func makeGallery(i int, includeScenes bool) *models.Gallery { } func createGalleries(ctx context.Context, n int) error { - gqb := db.Gallery - fqb := db.File + gqb := db.Gallery() + fqb := db.File() for i := 0; i < n; i++ { var fileIDs []models.FileID @@ -1573,7 +1588,7 @@ func getPerformerCustomFields(index int) map[string]interface{} { // createPerformers creates n performers with plain Name and o performers with camel cased NaMe included func createPerformers(ctx context.Context, n int, o int) error { - pqb := db.Performer + pqb := db.Performer() const namePlain = "Name" const nameNoCase = "NaMe" @@ -1699,7 +1714,7 @@ func getTagChildCount(id int) int { func tagStashID(i int) models.StashID { return models.StashID{ - StashID: getTagStringValue(i, "stashid"), + StashID: getUUID("stashid"), Endpoint: getTagStringValue(0, "endpoint"), } } @@ -1760,7 +1775,7 @@ func getStudioNullStringValue(index int, field string) string { return ret.String } -func createStudio(ctx context.Context, sqb *sqlite.StudioStore, name string, parentID *int) (*models.Studio, error) { +func createStudio(ctx context.Context, sqb database.StudioStore, name string, parentID *int) (*models.Studio, error) { studio := models.Studio{ Name: name, } @@ -1777,7 +1792,7 @@ func createStudio(ctx context.Context, sqb *sqlite.StudioStore, name string, par return &studio, nil } -func createStudioFromModel(ctx context.Context, sqb *sqlite.StudioStore, studio *models.Studio) error { +func createStudioFromModel(ctx context.Context, sqb database.StudioStore, studio *models.Studio) error { err := sqb.Create(ctx, studio) if err != nil { @@ -1812,7 +1827,7 @@ func getStudioStringList(index int, field string) []string { // createStudios creates n studios with plain Name and o studios with camel cased NaMe included func createStudios(ctx context.Context, n int, o int) error { - sqb := db.Studio + sqb := db.Studio() const namePlain = "Name" const nameNoCase = "NaMe" @@ -1991,7 +2006,7 @@ func linkGroupStudios(ctx context.Context, mqb models.GroupWriter) error { } func linkStudiosParent(ctx context.Context) error { - qb := db.Studio + qb := db.Studio() return doLinks(studioParentLinks, func(parentIndex, childIndex int) error { input := &models.StudioPartial{ ID: studioIDs[childIndex], diff --git a/pkg/sqlite/stash_id_test.go b/pkg/database/stash_id_test.go similarity index 69% rename from pkg/sqlite/stash_id_test.go rename to pkg/database/stash_id_test.go index a273c79609..ffcee27318 100644 --- a/pkg/sqlite/stash_id_test.go +++ b/pkg/database/stash_id_test.go @@ -1,7 +1,7 @@ -//go:build integration -// +build integration +//go:build db_integration +// +build db_integration -package sqlite_test +package database_test import ( "context" @@ -24,7 +24,7 @@ func testStashIDReaderWriter(ctx context.Context, t *testing.T, r stashIDReaderW testNoStashIDs(ctx, t, r, -1) // add stash ids - const stashIDStr = "stashID" + var stashIDStr = getUUID("stashID") const endpoint = "endpoint" stashID := models.StashID{ StashID: stashIDStr, @@ -39,11 +39,6 @@ func testStashIDReaderWriter(ctx context.Context, t *testing.T, r stashIDReaderW testStashIDs(ctx, t, r, id, []models.StashID{stashID}) - // update non-existing id - should return error - if err := r.UpdateStashIDs(ctx, -1, []models.StashID{stashID}); err == nil { - t.Error("expected error when updating non-existing id") - } - // remove stash ids and ensure was updated if err := r.UpdateStashIDs(ctx, id, []models.StashID{}); err != nil { t.Error(err.Error()) @@ -52,6 +47,35 @@ func testStashIDReaderWriter(ctx context.Context, t *testing.T, r stashIDReaderW testNoStashIDs(ctx, t, r, id) } +func testStashIDReaderWriterFail(ctx context.Context, t *testing.T, r stashIDReaderWriter, id int) { + // ensure no stash IDs to begin with + testNoStashIDs(ctx, t, r, id) + + // ensure GetStashIDs with non-existing also returns none + testNoStashIDs(ctx, t, r, -1) + + // add stash ids + var stashIDStr = getUUID("stashID") + const endpoint = "endpoint" + stashID := models.StashID{ + StashID: stashIDStr, + Endpoint: endpoint, + UpdatedAt: epochTime, + } + + // update stash ids and ensure was updated + if err := r.UpdateStashIDs(ctx, id, []models.StashID{stashID}); err != nil { + t.Error(err.Error()) + } + + testStashIDs(ctx, t, r, id, []models.StashID{stashID}) + + // update non-existing id - should return error + if err := r.UpdateStashIDs(ctx, -1, []models.StashID{stashID}); err == nil { + t.Error("expected error when updating non-existing id") + } +} + func testNoStashIDs(ctx context.Context, t *testing.T, r stashIDReaderWriter, id int) { t.Helper() stashIDs, err := r.GetStashIDs(ctx, id) diff --git a/pkg/database/studio.go b/pkg/database/studio.go new file mode 100644 index 0000000000..dad123f70b --- /dev/null +++ b/pkg/database/studio.go @@ -0,0 +1,40 @@ +package database + +import ( + "context" + + "github.com/stashapp/stash/pkg/models" +) + +type StudioStore interface { + All(ctx context.Context) ([]*models.Studio, error) + Count(ctx context.Context) (int, error) + Create(ctx context.Context, newObject *models.Studio) error + Destroy(ctx context.Context, id int) error + Find(ctx context.Context, id int) (*models.Studio, error) + FindByName(ctx context.Context, name string, nocase bool) (*models.Studio, error) + FindBySceneID(ctx context.Context, sceneID int) (*models.Studio, error) + FindByStashID(ctx context.Context, stashID models.StashID) ([]*models.Studio, error) + FindByStashIDStatus(ctx context.Context, hasStashID bool, stashboxEndpoint string) ([]*models.Studio, error) + FindChildren(ctx context.Context, id int) ([]*models.Studio, error) + FindMany(ctx context.Context, ids []int) ([]*models.Studio, error) + GetAliases(ctx context.Context, studioID int) ([]string, error) + GetImage(ctx context.Context, studioID int) ([]byte, error) + GetStashIDs(ctx context.Context, studioID int) ([]models.StashID, error) + HasImage(ctx context.Context, studioID int) (bool, error) + Query(ctx context.Context, studioFilter *models.StudioFilterType, findFilter *models.FindFilterType) ([]*models.Studio, int, error) + QueryCount(ctx context.Context, studioFilter *models.StudioFilterType, findFilter *models.FindFilterType) (int, error) + QueryForAutoTag(ctx context.Context, words []string) ([]*models.Studio, error) + Update(ctx context.Context, updatedObject *models.Studio) error + UpdateImage(ctx context.Context, studioID int, image []byte) error + UpdatePartial(ctx context.Context, input models.StudioPartial) (*models.Studio, error) + GetURLs(ctx context.Context, studioID int) ([]string, error) + + // blobJoinQueryBuilder + tagRelationshipStore +} + +type tagRelationshipStore interface { + CountByTagID(ctx context.Context, tagID int) (int, error) + GetTagIDs(ctx context.Context, id int) ([]int, error) +} diff --git a/pkg/sqlite/studio_test.go b/pkg/database/studio_test.go similarity index 93% rename from pkg/sqlite/studio_test.go rename to pkg/database/studio_test.go index 003877c779..0b1b760276 100644 --- a/pkg/sqlite/studio_test.go +++ b/pkg/database/studio_test.go @@ -1,7 +1,7 @@ -//go:build integration -// +build integration +//go:build db_integration +// +build db_integration -package sqlite_test +package database_test import ( "context" @@ -18,7 +18,7 @@ import ( func TestStudioFindByName(t *testing.T) { withTxn(func(ctx context.Context) error { - sqb := db.Studio + sqb := db.Studio() name := studioNames[studioIdxWithScene] // find a studio by name @@ -70,7 +70,7 @@ func TestStudioQueryNameOr(t *testing.T) { } withTxn(func(ctx context.Context) error { - sqb := db.Studio + sqb := db.Studio() studios := queryStudio(ctx, t, sqb, &studioFilter, nil) @@ -83,7 +83,7 @@ func TestStudioQueryNameOr(t *testing.T) { } func loadStudioRelationships(ctx context.Context, t *testing.T, s *models.Studio) error { - if err := s.LoadURLs(ctx, db.Studio); err != nil { + if err := s.LoadURLs(ctx, db.Studio()); err != nil { return err } @@ -111,7 +111,7 @@ func TestStudioQueryNameAndUrl(t *testing.T) { } withTxn(func(ctx context.Context) error { - sqb := db.Studio + sqb := db.Studio() studios := queryStudio(ctx, t, sqb, &studioFilter, nil) @@ -119,7 +119,7 @@ func TestStudioQueryNameAndUrl(t *testing.T) { return nil } - if err := studios[0].LoadURLs(ctx, db.Studio); err != nil { + if err := studios[0].LoadURLs(ctx, db.Studio()); err != nil { t.Errorf("Error loading studio relationships: %v", err) } @@ -155,12 +155,12 @@ func TestStudioQueryNameNotUrl(t *testing.T) { } withTxn(func(ctx context.Context) error { - sqb := db.Studio + sqb := db.Studio() studios := queryStudio(ctx, t, sqb, &studioFilter, nil) for _, studio := range studios { - if err := studio.LoadURLs(ctx, db.Studio); err != nil { + if err := studio.LoadURLs(ctx, db.Studio()); err != nil { t.Errorf("Error loading studio relationships: %v", err) } @@ -192,7 +192,7 @@ func TestStudioIllegalQuery(t *testing.T) { } withTxn(func(ctx context.Context) error { - sqb := db.Studio + sqb := db.Studio() _, _, err := sqb.Query(ctx, studioFilter, nil) assert.NotNil(err) @@ -218,7 +218,7 @@ func TestStudioQueryIgnoreAutoTag(t *testing.T) { IgnoreAutoTag: &ignoreAutoTag, } - sqb := db.Studio + sqb := db.Studio() studios := queryStudio(ctx, t, sqb, &studioFilter, nil) @@ -233,7 +233,7 @@ func TestStudioQueryIgnoreAutoTag(t *testing.T) { func TestStudioQueryForAutoTag(t *testing.T) { withTxn(func(ctx context.Context) error { - tqb := db.Studio + tqb := db.Studio() name := studioNames[studioIdxWithGroup] // find a studio by name @@ -261,7 +261,7 @@ func TestStudioQueryForAutoTag(t *testing.T) { func TestStudioQueryParent(t *testing.T) { withTxn(func(ctx context.Context) error { - sqb := db.Studio + sqb := db.Studio() studioCriterion := models.MultiCriterionInput{ Value: []string{ strconv.Itoa(studioIDs[studioIdxWithChildStudio]), @@ -311,18 +311,18 @@ func TestStudioDestroyParent(t *testing.T) { // create parent and child studios if err := withTxn(func(ctx context.Context) error { - createdParent, err := createStudio(ctx, db.Studio, parentName, nil) + createdParent, err := createStudio(ctx, db.Studio(), parentName, nil) if err != nil { return fmt.Errorf("Error creating parent studio: %s", err.Error()) } parentID := createdParent.ID - createdChild, err := createStudio(ctx, db.Studio, childName, &parentID) + createdChild, err := createStudio(ctx, db.Studio(), childName, &parentID) if err != nil { return fmt.Errorf("Error creating child studio: %s", err.Error()) } - sqb := db.Studio + sqb := db.Studio() // destroy the parent err = sqb.Destroy(ctx, createdParent.ID) @@ -344,7 +344,7 @@ func TestStudioDestroyParent(t *testing.T) { func TestStudioFindChildren(t *testing.T) { withTxn(func(ctx context.Context) error { - sqb := db.Studio + sqb := db.Studio() studios, err := sqb.FindChildren(ctx, studioIDs[studioIdxWithChildStudio]) @@ -373,18 +373,18 @@ func TestStudioUpdateClearParent(t *testing.T) { // create parent and child studios if err := withTxn(func(ctx context.Context) error { - createdParent, err := createStudio(ctx, db.Studio, parentName, nil) + createdParent, err := createStudio(ctx, db.Studio(), parentName, nil) if err != nil { return fmt.Errorf("Error creating parent studio: %s", err.Error()) } parentID := createdParent.ID - createdChild, err := createStudio(ctx, db.Studio, childName, &parentID) + createdChild, err := createStudio(ctx, db.Studio(), childName, &parentID) if err != nil { return fmt.Errorf("Error creating child studio: %s", err.Error()) } - sqb := db.Studio + sqb := db.Studio() // clear the parent id from the child input := models.StudioPartial{ @@ -410,11 +410,11 @@ func TestStudioUpdateClearParent(t *testing.T) { func TestStudioUpdateStudioImage(t *testing.T) { if err := withTxn(func(ctx context.Context) error { - qb := db.Studio + qb := db.Studio() // create studio to test against const name = "TestStudioUpdateStudioImage" - created, err := createStudio(ctx, db.Studio, name, nil) + created, err := createStudio(ctx, db.Studio(), name, nil) if err != nil { return fmt.Errorf("Error creating studio: %s", err.Error()) } @@ -446,7 +446,7 @@ func TestStudioQuerySceneCount(t *testing.T) { func verifyStudiosSceneCount(t *testing.T, sceneCountCriterion models.IntCriterionInput) { withTxn(func(ctx context.Context) error { - sqb := db.Studio + sqb := db.Studio() studioFilter := models.StudioFilterType{ SceneCount: &sceneCountCriterion, } @@ -455,7 +455,7 @@ func verifyStudiosSceneCount(t *testing.T, sceneCountCriterion models.IntCriteri assert.Greater(t, len(studios), 0) for _, studio := range studios { - sceneCount, err := db.Scene.CountByStudioID(ctx, studio.ID) + sceneCount, err := db.Scene().CountByStudioID(ctx, studio.ID) if err != nil { return err } @@ -487,7 +487,7 @@ func TestStudioQueryImageCount(t *testing.T) { func verifyStudiosImageCount(t *testing.T, imageCountCriterion models.IntCriterionInput) { withTxn(func(ctx context.Context) error { - sqb := db.Studio + sqb := db.Studio() studioFilter := models.StudioFilterType{ ImageCount: &imageCountCriterion, } @@ -498,7 +498,7 @@ func verifyStudiosImageCount(t *testing.T, imageCountCriterion models.IntCriteri for _, studio := range studios { pp := 0 - result, err := db.Image.Query(ctx, models.ImageQueryOptions{ + result, err := db.Image().Query(ctx, models.ImageQueryOptions{ QueryOptions: models.QueryOptions{ FindFilter: &models.FindFilterType{ PerPage: &pp, @@ -543,7 +543,7 @@ func TestStudioQueryGalleryCount(t *testing.T) { func verifyStudiosGalleryCount(t *testing.T, galleryCountCriterion models.IntCriterionInput) { withTxn(func(ctx context.Context) error { - sqb := db.Studio + sqb := db.Studio() studioFilter := models.StudioFilterType{ GalleryCount: &galleryCountCriterion, } @@ -554,7 +554,7 @@ func verifyStudiosGalleryCount(t *testing.T, galleryCountCriterion models.IntCri for _, studio := range studios { pp := 0 - _, count, err := db.Gallery.Query(ctx, &models.GalleryFilterType{ + _, count, err := db.Gallery().Query(ctx, &models.GalleryFilterType{ Studios: &models.HierarchicalMultiCriterionInput{ Value: []string{strconv.Itoa(studio.ID)}, Modifier: models.CriterionModifierIncludes, @@ -574,11 +574,11 @@ func verifyStudiosGalleryCount(t *testing.T, galleryCountCriterion models.IntCri func TestStudioStashIDs(t *testing.T) { if err := withRollbackTxn(func(ctx context.Context) error { - qb := db.Studio + qb := db.Studio() // create studio to test against const name = "TestStudioStashIDs" - created, err := createStudio(ctx, db.Studio, name, nil) + created, err := createStudio(ctx, db.Studio(), name, nil) if err != nil { return fmt.Errorf("Error creating studio: %s", err.Error()) } @@ -600,7 +600,7 @@ func TestStudioStashIDs(t *testing.T) { } func testStudioStashIDs(ctx context.Context, t *testing.T, s *models.Studio) { - qb := db.Studio + qb := db.Studio() if err := s.LoadStashIDs(ctx, qb); err != nil { t.Error(err.Error()) @@ -611,7 +611,7 @@ func testStudioStashIDs(ctx context.Context, t *testing.T, s *models.Studio) { assert.Len(t, s.StashIDs.List(), 0) // add stash ids - const stashIDStr = "stashID" + var stashIDStr = getUUID("stashID") const endpoint = "endpoint" stashID := models.StashID{ StashID: stashIDStr, @@ -678,7 +678,7 @@ func TestStudioQueryURL(t *testing.T) { verifyFn := func(ctx context.Context, g *models.Studio) { t.Helper() - if err := g.LoadURLs(ctx, db.Studio); err != nil { + if err := g.LoadURLs(ctx, db.Studio()); err != nil { t.Errorf("Error loading studio relationships: %v", err) return } @@ -732,7 +732,7 @@ func TestStudioQueryRating(t *testing.T) { func queryStudios(ctx context.Context, t *testing.T, studioFilter *models.StudioFilterType, findFilter *models.FindFilterType) []*models.Studio { t.Helper() - studios, _, err := db.Studio.Query(ctx, studioFilter, findFilter) + studios, _, err := db.Studio().Query(ctx, studioFilter, findFilter) if err != nil { t.Errorf("Error querying studio: %s", err.Error()) } @@ -814,7 +814,7 @@ func TestStudioQueryTagCount(t *testing.T) { func verifyStudiosTagCount(t *testing.T, tagCountCriterion models.IntCriterionInput) { withTxn(func(ctx context.Context) error { - sqb := db.Studio + sqb := db.Studio() studioFilter := models.StudioFilterType{ TagCount: &tagCountCriterion, } @@ -837,7 +837,7 @@ func verifyStudiosTagCount(t *testing.T, tagCountCriterion models.IntCriterionIn func verifyStudioQuery(t *testing.T, filter models.StudioFilterType, verifyFn func(ctx context.Context, s *models.Studio)) { withTxn(func(ctx context.Context) error { t.Helper() - sqb := db.Studio + sqb := db.Studio() studios := queryStudio(ctx, t, sqb, &filter, nil) @@ -854,7 +854,7 @@ func verifyStudioQuery(t *testing.T, filter models.StudioFilterType, verifyFn fu func verifyStudiosRating(t *testing.T, ratingCriterion models.IntCriterionInput) { withTxn(func(ctx context.Context) error { - sqb := db.Studio + sqb := db.Studio() studioFilter := models.StudioFilterType{ Rating100: &ratingCriterion, } @@ -875,7 +875,7 @@ func verifyStudiosRating(t *testing.T, ratingCriterion models.IntCriterionInput) func TestStudioQueryIsMissingRating(t *testing.T) { withTxn(func(ctx context.Context) error { - sqb := db.Studio + sqb := db.Studio() isMissing := "rating" studioFilter := models.StudioFilterType{ IsMissing: &isMissing, @@ -951,7 +951,7 @@ func TestStudioQueryAlias(t *testing.T) { verifyFn := func(ctx context.Context, studio *models.Studio) { t.Helper() - aliases, err := db.Studio.GetAliases(ctx, studio.ID) + aliases, err := db.Studio().GetAliases(ctx, studio.ID) if err != nil { t.Errorf("Error querying studios: %s", err.Error()) } @@ -986,11 +986,11 @@ func TestStudioQueryAlias(t *testing.T) { func TestStudioAlias(t *testing.T) { if err := withRollbackTxn(func(ctx context.Context) error { - qb := db.Studio + qb := db.Studio() // create studio to test against const name = "TestStudioAlias" - created, err := createStudio(ctx, db.Studio, name, nil) + created, err := createStudio(ctx, db.Studio(), name, nil) if err != nil { return fmt.Errorf("Error creating studio: %s", err.Error()) } @@ -1012,7 +1012,7 @@ func TestStudioAlias(t *testing.T) { } func testStudioAlias(ctx context.Context, t *testing.T, s *models.Studio) { - qb := db.Studio + qb := db.Studio() if err := s.LoadAliases(ctx, qb); err != nil { t.Error(err.Error()) return @@ -1070,13 +1070,14 @@ func TestStudioQueryFast(t *testing.T) { tsString := "test" tsInt := 1 + tsId := "1" testStringCriterion := models.StringCriterionInput{ Value: tsString, Modifier: models.CriterionModifierEquals, } - testIncludesMultiCriterion := models.MultiCriterionInput{ - Value: []string{tsString}, + testIncludesMultiCriterionId := models.MultiCriterionInput{ + Value: []string{tsId}, Modifier: models.CriterionModifierIncludes, } testIntCriterion := models.IntCriterionInput{ @@ -1106,7 +1107,7 @@ func TestStudioQueryFast(t *testing.T) { SceneCount: &testIntCriterion, } parentsFilter := models.StudioFilterType{ - Parents: &testIncludesMultiCriterion, + Parents: &testIncludesMultiCriterionId, } filters := []models.StudioFilterType{nameFilter, aliasesFilter, stashIDFilter, urlFilter, ratingFilter, sceneCountFilter, imageCountFilter, parentsFilter} @@ -1134,7 +1135,7 @@ func TestStudioQueryFast(t *testing.T) { } withTxn(func(ctx context.Context) error { - sqb := db.Studio + sqb := db.Studio() for _, f := range filters { for _, ff := range findFilters { _, _, err := sqb.Query(ctx, &f, &ff) diff --git a/pkg/database/tag.go b/pkg/database/tag.go new file mode 100644 index 0000000000..ba4755149e --- /dev/null +++ b/pkg/database/tag.go @@ -0,0 +1,51 @@ +package database + +import ( + "context" + + "github.com/stashapp/stash/pkg/models" +) + +type TagStore interface { + All(ctx context.Context) ([]*models.Tag, error) + Count(ctx context.Context) (int, error) + CountByChildTagID(ctx context.Context, childID int) (int, error) + CountByParentTagID(ctx context.Context, parentID int) (int, error) + Create(ctx context.Context, newObject *models.Tag) error + Destroy(ctx context.Context, id int) error + Find(ctx context.Context, id int) (*models.Tag, error) + FindAllAncestors(ctx context.Context, tagID int, excludeIDs []int) ([]*models.TagPath, error) + FindAllDescendants(ctx context.Context, tagID int, excludeIDs []int) ([]*models.TagPath, error) + FindByChildTagID(ctx context.Context, parentID int) ([]*models.Tag, error) + FindByGalleryID(ctx context.Context, galleryID int) ([]*models.Tag, error) + FindByGroupID(ctx context.Context, groupID int) ([]*models.Tag, error) + FindByImageID(ctx context.Context, imageID int) ([]*models.Tag, error) + FindByName(ctx context.Context, name string, nocase bool) (*models.Tag, error) + FindByNames(ctx context.Context, names []string, nocase bool) ([]*models.Tag, error) + FindByParentTagID(ctx context.Context, parentID int) ([]*models.Tag, error) + FindByPerformerID(ctx context.Context, performerID int) ([]*models.Tag, error) + FindBySceneID(ctx context.Context, sceneID int) ([]*models.Tag, error) + FindBySceneMarkerID(ctx context.Context, sceneMarkerID int) ([]*models.Tag, error) + FindByStudioID(ctx context.Context, studioID int) ([]*models.Tag, error) + FindByStashID(ctx context.Context, stashID models.StashID) ([]*models.Tag, error) + FindMany(ctx context.Context, ids []int) ([]*models.Tag, error) + GetAliases(ctx context.Context, tagID int) ([]string, error) + GetChildIDs(ctx context.Context, relatedID int) ([]int, error) + GetImage(ctx context.Context, tagID int) ([]byte, error) + GetParentIDs(ctx context.Context, relatedID int) ([]int, error) + HasImage(ctx context.Context, tagID int) (bool, error) + Merge(ctx context.Context, source []int, destination int) error + Query(ctx context.Context, tagFilter *models.TagFilterType, findFilter *models.FindFilterType) ([]*models.Tag, int, error) + QueryForAutoTag(ctx context.Context, words []string) ([]*models.Tag, error) + Update(ctx context.Context, updatedObject *models.Tag) error + UpdateAliases(ctx context.Context, tagID int, aliases []string) error + UpdateChildTags(ctx context.Context, tagID int, childIDs []int) error + UpdateImage(ctx context.Context, tagID int, image []byte) error + UpdateParentTags(ctx context.Context, tagID int, parentIDs []int) error + UpdatePartial(ctx context.Context, id int, partial models.TagPartial) (*models.Tag, error) + + GetStashIDs(ctx context.Context, tagID int) ([]models.StashID, error) + UpdateStashIDs(ctx context.Context, tagID int, stashIDs []models.StashID) error + + // blobJoinQueryBuilder +} diff --git a/pkg/sqlite/tag_test.go b/pkg/database/tag_test.go similarity index 95% rename from pkg/sqlite/tag_test.go rename to pkg/database/tag_test.go index f1bac19b24..a30a67b530 100644 --- a/pkg/sqlite/tag_test.go +++ b/pkg/database/tag_test.go @@ -1,7 +1,7 @@ -//go:build integration -// +build integration +//go:build db_integration +// +build db_integration -package sqlite_test +package database_test import ( "context" @@ -17,7 +17,7 @@ import ( func TestMarkerFindBySceneMarkerID(t *testing.T) { withTxn(func(ctx context.Context) error { - tqb := db.Tag + tqb := db.Tag() markerID := markerIDs[markerIdxWithTag] @@ -44,7 +44,7 @@ func TestMarkerFindBySceneMarkerID(t *testing.T) { func TestTagFindByGroupID(t *testing.T) { withTxn(func(ctx context.Context) error { - tqb := db.Tag + tqb := db.Tag() groupID := groupIDs[groupIdxWithTag] @@ -71,7 +71,7 @@ func TestTagFindByGroupID(t *testing.T) { func TestTagFindByName(t *testing.T) { withTxn(func(ctx context.Context) error { - tqb := db.Tag + tqb := db.Tag() name := tagNames[tagIdxWithScene] // find a tag by name @@ -107,7 +107,7 @@ func TestTagQueryIgnoreAutoTag(t *testing.T) { IgnoreAutoTag: &ignoreAutoTag, } - sqb := db.Tag + sqb := db.Tag() tags := queryTags(ctx, t, sqb, &tagFilter, nil) @@ -122,7 +122,7 @@ func TestTagQueryIgnoreAutoTag(t *testing.T) { func TestTagQueryForAutoTag(t *testing.T) { withTxn(func(ctx context.Context) error { - tqb := db.Tag + tqb := db.Tag() name := tagNames[tagIdx1WithScene] // find a tag by name @@ -156,7 +156,7 @@ func TestTagFindByNames(t *testing.T) { var names []string withTxn(func(ctx context.Context) error { - tqb := db.Tag + tqb := db.Tag() names = append(names, tagNames[tagIdxWithScene]) // find tags by names @@ -201,7 +201,7 @@ func TestTagFindByNames(t *testing.T) { func TestTagQuerySort(t *testing.T) { withTxn(func(ctx context.Context) error { - sqb := db.Tag + sqb := db.Tag() sortBy := "scenes_count" dir := models.SortDirectionEnumDesc @@ -286,7 +286,7 @@ func TestTagQueryAlias(t *testing.T) { } verifyFn := func(ctx context.Context, tag *models.Tag) { - aliases, err := db.Tag.GetAliases(ctx, tag.ID) + aliases, err := db.Tag().GetAliases(ctx, tag.ID) if err != nil { t.Errorf("Error querying tags: %s", err.Error()) } @@ -321,7 +321,7 @@ func TestTagQueryAlias(t *testing.T) { func verifyTagQuery(t *testing.T, tagFilter *models.TagFilterType, findFilter *models.FindFilterType, verifyFn func(ctx context.Context, t *models.Tag)) { withTxn(func(ctx context.Context) error { - sqb := db.Tag + sqb := db.Tag() tags := queryTags(ctx, t, sqb, tagFilter, findFilter) @@ -482,7 +482,7 @@ func TestTagQuery(t *testing.T) { runWithRollbackTxn(t, tt.name, func(t *testing.T, ctx context.Context) { assert := assert.New(t) - tags, _, err := db.Tag.Query(ctx, tt.filter, tt.findFilter) + tags, _, err := db.Tag().Query(ctx, tt.filter, tt.findFilter) if (err != nil) != tt.wantErr { t.Errorf("PerformerStore.Query() error = %v, wantErr %v", err, tt.wantErr) return @@ -504,7 +504,7 @@ func TestTagQuery(t *testing.T) { func TestTagQueryIsMissingImage(t *testing.T) { withTxn(func(ctx context.Context) error { - qb := db.Tag + qb := db.Tag() isMissing := "image" tagFilter := models.TagFilterType{ IsMissing: &isMissing, @@ -558,7 +558,7 @@ func TestTagQuerySceneCount(t *testing.T) { func verifyTagSceneCount(t *testing.T, sceneCountCriterion models.IntCriterionInput) { withTxn(func(ctx context.Context) error { - qb := db.Tag + qb := db.Tag() tagFilter := models.TagFilterType{ SceneCount: &sceneCountCriterion, } @@ -597,7 +597,7 @@ func TestTagQueryMarkerCount(t *testing.T) { func verifyTagMarkerCount(t *testing.T, markerCountCriterion models.IntCriterionInput) { withTxn(func(ctx context.Context) error { - qb := db.Tag + qb := db.Tag() tagFilter := models.TagFilterType{ MarkerCount: &markerCountCriterion, } @@ -636,7 +636,7 @@ func TestTagQueryImageCount(t *testing.T) { func verifyTagImageCount(t *testing.T, imageCountCriterion models.IntCriterionInput) { withTxn(func(ctx context.Context) error { - qb := db.Tag + qb := db.Tag() tagFilter := models.TagFilterType{ ImageCount: &imageCountCriterion, } @@ -675,7 +675,7 @@ func TestTagQueryGalleryCount(t *testing.T) { func verifyTagGalleryCount(t *testing.T, imageCountCriterion models.IntCriterionInput) { withTxn(func(ctx context.Context) error { - qb := db.Tag + qb := db.Tag() tagFilter := models.TagFilterType{ GalleryCount: &imageCountCriterion, } @@ -714,7 +714,7 @@ func TestTagQueryPerformerCount(t *testing.T) { func verifyTagPerformerCount(t *testing.T, imageCountCriterion models.IntCriterionInput) { withTxn(func(ctx context.Context) error { - qb := db.Tag + qb := db.Tag() tagFilter := models.TagFilterType{ PerformerCount: &imageCountCriterion, } @@ -753,7 +753,7 @@ func TestTagQueryStudioCount(t *testing.T) { func verifyTagStudioCount(t *testing.T, imageCountCriterion models.IntCriterionInput) { withTxn(func(ctx context.Context) error { - qb := db.Tag + qb := db.Tag() tagFilter := models.TagFilterType{ StudioCount: &imageCountCriterion, } @@ -792,7 +792,7 @@ func TestTagQueryParentCount(t *testing.T) { func verifyTagParentCount(t *testing.T, sceneCountCriterion models.IntCriterionInput) { withTxn(func(ctx context.Context) error { - qb := db.Tag + qb := db.Tag() tagFilter := models.TagFilterType{ ParentCount: &sceneCountCriterion, } @@ -832,7 +832,7 @@ func TestTagQueryChildCount(t *testing.T) { func verifyTagChildCount(t *testing.T, sceneCountCriterion models.IntCriterionInput) { withTxn(func(ctx context.Context) error { - qb := db.Tag + qb := db.Tag() tagFilter := models.TagFilterType{ ChildCount: &sceneCountCriterion, } @@ -854,7 +854,7 @@ func verifyTagChildCount(t *testing.T, sceneCountCriterion models.IntCriterionIn func TestTagQueryParent(t *testing.T) { withTxn(func(ctx context.Context) error { const nameField = "Name" - sqb := db.Tag + sqb := db.Tag() tagCriterion := models.HierarchicalMultiCriterionInput{ Value: []string{ strconv.Itoa(tagIDs[tagIdxWithChildTag]), @@ -932,7 +932,7 @@ func TestTagQueryChild(t *testing.T) { withTxn(func(ctx context.Context) error { const nameField = "Name" - sqb := db.Tag + sqb := db.Tag() tagCriterion := models.HierarchicalMultiCriterionInput{ Value: []string{ strconv.Itoa(tagIDs[tagIdxWithParentTag]), @@ -1008,7 +1008,7 @@ func TestTagQueryChild(t *testing.T) { func TestTagUpdateTagImage(t *testing.T) { if err := withTxn(func(ctx context.Context) error { - qb := db.Tag + qb := db.Tag() // create tag to test against const name = "TestTagUpdateTagImage" @@ -1028,7 +1028,7 @@ func TestTagUpdateTagImage(t *testing.T) { func TestTagUpdateAlias(t *testing.T) { if err := withTxn(func(ctx context.Context) error { - qb := db.Tag + qb := db.Tag() // create tag to test against const name = "TestTagUpdateAlias" @@ -1061,7 +1061,7 @@ func TestTagUpdateAlias(t *testing.T) { func TestTagStashIDs(t *testing.T) { if err := withTxn(func(ctx context.Context) error { - qb := db.Tag + qb := db.Tag() // create tag to test against const name = "TestTagStashIDs" @@ -1083,7 +1083,7 @@ func TestTagStashIDs(t *testing.T) { func TestTagFindByStashID(t *testing.T) { withTxn(func(ctx context.Context) error { - qb := db.Tag + qb := db.Tag() // create tag to test against const name = "TestTagFindByStashID" @@ -1124,8 +1124,8 @@ func TestTagMerge(t *testing.T) { // merge tests - perform these in a transaction that we'll rollback if err := withRollbackTxn(func(ctx context.Context) error { - qb := db.Tag - mqb := db.SceneMarker + qb := db.Tag() + mqb := db.SceneMarker() // try merging into same tag err := qb.Merge(ctx, []int{tagIDs[tagIdx1WithScene]}, tagIDs[tagIdx1WithScene]) @@ -1183,11 +1183,11 @@ func TestTagMerge(t *testing.T) { } // ensure scene points to new tag - s, err := db.Scene.Find(ctx, sceneIDs[sceneIdxWithTwoTags]) + s, err := db.Scene().Find(ctx, sceneIDs[sceneIdxWithTwoTags]) if err != nil { return err } - if err := s.LoadTagIDs(ctx, db.Scene); err != nil { + if err := s.LoadTagIDs(ctx, db.Scene()); err != nil { return err } sceneTagIDs := s.TagIDs.List() @@ -1210,19 +1210,19 @@ func TestTagMerge(t *testing.T) { assert.Contains(markerTagIDs, destID) // ensure image points to new tag - imageTagIDs, err := db.Image.GetTagIDs(ctx, imageIDs[imageIdxWithTwoTags]) + imageTagIDs, err := db.Image().GetTagIDs(ctx, imageIDs[imageIdxWithTwoTags]) if err != nil { return err } assert.Contains(imageTagIDs, destID) - g, err := db.Gallery.Find(ctx, galleryIDs[galleryIdxWithTwoTags]) + g, err := db.Gallery().Find(ctx, galleryIDs[galleryIdxWithTwoTags]) if err != nil { return err } - if err := g.LoadTagIDs(ctx, db.Gallery); err != nil { + if err := g.LoadTagIDs(ctx, db.Gallery()); err != nil { return err } @@ -1230,7 +1230,7 @@ func TestTagMerge(t *testing.T) { assert.Contains(g.TagIDs.List(), destID) // ensure performer points to new tag - performerTagIDs, err := db.Performer.GetTagIDs(ctx, performerIDs[performerIdxWithTwoTags]) + performerTagIDs, err := db.Performer().GetTagIDs(ctx, performerIDs[performerIdxWithTwoTags]) if err != nil { return err } @@ -1238,7 +1238,7 @@ func TestTagMerge(t *testing.T) { assert.Contains(performerTagIDs, destID) // ensure studio points to new tag - studioTagIDs, err := db.Studio.GetTagIDs(ctx, studioIDs[studioIdxWithTwoTags]) + studioTagIDs, err := db.Studio().GetTagIDs(ctx, studioIDs[studioIdxWithTwoTags]) if err != nil { return err } @@ -1246,11 +1246,11 @@ func TestTagMerge(t *testing.T) { assert.Contains(studioTagIDs, destID) // ensure group points to new tag - group, err := db.Group.Find(ctx, groupIDs[groupIdxWithTwoTags]) + group, err := db.Group().Find(ctx, groupIDs[groupIdxWithTwoTags]) if err != nil { return err } - if err := group.LoadTagIDs(ctx, db.Group); err != nil { + if err := group.LoadTagIDs(ctx, db.Group()); err != nil { return err } groupTagIDs := group.TagIDs.List() diff --git a/pkg/sqlite/transaction_test.go b/pkg/database/transaction_test.go similarity index 87% rename from pkg/sqlite/transaction_test.go rename to pkg/database/transaction_test.go index 070f8b2c9a..4a32b900ea 100644 --- a/pkg/sqlite/transaction_test.go +++ b/pkg/database/transaction_test.go @@ -1,7 +1,7 @@ -//go:build integration -// +build integration +//go:build db_integration +// +build db_integration -package sqlite_test +package database_test import ( "context" @@ -36,11 +36,11 @@ import ( // Title: "test", // } -// if err := db.Scene.Create(ctx, scene, nil); err != nil { +// if err := db.Scene().Create(ctx, scene, nil); err != nil { // return err // } -// if err := db.Scene.Destroy(ctx, scene.ID); err != nil { +// if err := db.Scene().Destroy(ctx, scene.ID); err != nil { // return err // } // } @@ -94,7 +94,7 @@ func waitForOtherThread(c chan struct{}) error { // Title: "test", // } -// if err := db.Scene.Create(ctx, scene, nil); err != nil { +// if err := db.Scene().Create(ctx, scene, nil); err != nil { // return err // } @@ -106,7 +106,7 @@ func waitForOtherThread(c chan struct{}) error { // return err // } -// if err := db.Scene.Destroy(ctx, scene.ID); err != nil { +// if err := db.Scene().Destroy(ctx, scene.ID); err != nil { // return err // } @@ -139,7 +139,7 @@ func waitForOtherThread(c chan struct{}) error { // // expect error when we try to do this, as the other thread has already // // modified this table // // this takes time to fail, so we need to wait for it -// if err := db.Scene.Create(ctx, scene, nil); err != nil { +// if err := db.Scene().Create(ctx, scene, nil); err != nil { // if !db.IsLocked(err) { // t.Errorf("unexpected error: %v", err) // } @@ -169,7 +169,7 @@ func TestConcurrentExclusiveAndReadTxn(t *testing.T) { Title: "test", } - if err := db.Scene.Create(ctx, scene, nil); err != nil { + if err := db.Scene().Create(ctx, scene, nil); err != nil { return err } @@ -181,7 +181,7 @@ func TestConcurrentExclusiveAndReadTxn(t *testing.T) { return err } - if err := db.Scene.Destroy(ctx, scene.ID); err != nil { + if err := db.Scene().Destroy(ctx, scene.ID); err != nil { return err } @@ -207,7 +207,7 @@ func TestConcurrentExclusiveAndReadTxn(t *testing.T) { } }() - if _, err := db.Scene.Find(ctx, sceneIDs[sceneIdx1WithPerformer]); err != nil { + if _, err := db.Scene().Find(ctx, sceneIDs[sceneIdx1WithPerformer]); err != nil { t.Errorf("unexpected error: %v", err) return err } @@ -241,11 +241,11 @@ func TestConcurrentExclusiveAndReadTxn(t *testing.T) { // Title: "test", // } -// if err := db.Scene.Create(ctx, scene, nil); err != nil { +// if err := db.Scene().Create(ctx, scene, nil); err != nil { // return err // } -// if err := db.Scene.Destroy(ctx, scene.ID); err != nil { +// if err := db.Scene().Destroy(ctx, scene.ID); err != nil { // return err // } // } @@ -267,7 +267,7 @@ func TestConcurrentExclusiveAndReadTxn(t *testing.T) { // for l := 0; l < loops; l++ { // if err := txn.WithReadTxn(ctx, db, func(ctx context.Context) error { // for ll := 0; ll < innerLoops; ll++ { -// if _, err := db.Scene.Find(ctx, sceneIDs[ll%totalScenes]); err != nil { +// if _, err := db.Scene().Find(ctx, sceneIDs[ll%totalScenes]); err != nil { // return err // } // } diff --git a/pkg/models/find_filter.go b/pkg/models/find_filter.go index 9934a9ea9c..e9f5a1feaf 100644 --- a/pkg/models/find_filter.go +++ b/pkg/models/find_filter.go @@ -118,6 +118,10 @@ func (ff FindFilterType) IsGetAll() bool { return ff.PerPage != nil && *ff.PerPage < 0 } +func (ff FindFilterType) IsCounting() bool { + return ff.PerPage != nil && *ff.PerPage == 0 +} + // BatchFindFilter returns a FindFilterType suitable for batch finding // using the provided batch size. func BatchFindFilter(batchSize int) *FindFilterType { diff --git a/pkg/models/relationships.go b/pkg/models/relationships.go index 5495f858b1..a899490ec5 100644 --- a/pkg/models/relationships.go +++ b/pkg/models/relationships.go @@ -2,6 +2,7 @@ package models import ( "context" + "slices" "github.com/stashapp/stash/pkg/sliceutil" ) @@ -86,6 +87,10 @@ func (r RelatedIDs) Loaded() bool { return r.list != nil } +func (r RelatedIDs) Sort() { + slices.Sort(r.list) +} + func (r RelatedIDs) mustLoaded() { if !r.Loaded() { panic("list has not been loaded") diff --git a/pkg/postgres/anonymise.go b/pkg/postgres/anonymise.go new file mode 100644 index 0000000000..2b9d55212c --- /dev/null +++ b/pkg/postgres/anonymise.go @@ -0,0 +1,145 @@ +package postgres + +import ( + "context" + "fmt" + + "github.com/doug-martin/goqu/v9" + "github.com/doug-martin/goqu/v9/exp" + "github.com/jmoiron/sqlx" + "github.com/stashapp/stash/pkg/sqlite" + "github.com/stashapp/stash/pkg/txn" +) + +const ( + batchSize = 5000 +) + +type Anonymiser struct { + *sqlite.Database + sourceDB *Database +} + +func NewAnonymiser(db *Database, outPath string) (*sqlite.Anonymiser, error) { + newDB := &Anonymiser{Database: sqlite.NewDatabase(), sourceDB: db} + if err := newDB.Open(outPath); err != nil { + return nil, fmt.Errorf("opening %s: %w", outPath, err) + } + + return sqlite.PassAnonymiser(newDB) +} + +func (db *Anonymiser) GetSqliteDatabase() *sqlite.Database { + return db.Database +} + +func (db *Anonymiser) FetchAll(ctx context.Context) error { + var sqlite_dialect = goqu.Dialect("sqlite3") + + ctx, err := db.Begin(ctx, true) + if err != nil { + return fmt.Errorf("begin tx: %w", err) + } + + for _, table := range []exp.IdentifierExpression{ + goqu.I(fileTable), + goqu.I(fingerprintTable), + goqu.I(folderTable), + goqu.I(galleryTable), + goqu.I(galleriesChaptersTable), + goqu.I(galleriesFilesTable), + goqu.I(galleriesImagesTable), + goqu.I(galleriesTagsTable), + goqu.I(galleriesURLsTable), + goqu.I(groupURLsTable), + goqu.I(groupTable), + goqu.I(groupRelationsTable), + goqu.I(groupsScenesTable), + goqu.I(groupsTagsTable), + goqu.I(imageFileTable), + goqu.I(imagesURLsTable), + goqu.I(imageTable), + goqu.I(imagesFilesTable), + goqu.I(imagesTagsTable), + goqu.I(performersAliasesTable), + goqu.I("performer_stash_ids"), + goqu.I(performerURLsTable), + goqu.I(performerTable), + goqu.I(performersGalleriesTable), + goqu.I(performersImagesTable), + goqu.I(performersScenesTable), + goqu.I(performersTagsTable), + goqu.I(savedFilterTable), + goqu.I(sceneMarkerTable), + goqu.I("scene_markers_tags"), + goqu.I(scenesURLsTable), + goqu.I(sceneTable), + goqu.I(scenesFilesTable), + goqu.I(scenesGalleriesTable), + goqu.I(scenesODatesTable), + goqu.I(scenesTagsTable), + goqu.I(scenesViewDatesTable), + goqu.I(studioAliasesTable), + goqu.I("studio_stash_ids"), + goqu.I(studioTable), + goqu.I(studiosTagsTable), + goqu.I(tagAliasesTable), + goqu.I(tagTable), + goqu.I(tagRelationsTable), + goqu.I("tag_stash_ids"), + goqu.I(videoCaptionsTable), + goqu.I(videoFileTable), + } { + offset := 0 + for { + q := dialect.From(table).Select(table.All()).Limit(uint(batchSize)).Offset(uint(offset)) + var rowsSlice []map[string]interface{} + + // Fetch + if err := txn.WithTxn(ctx, db.sourceDB, func(ctx context.Context) error { + if err := queryFunc(ctx, q, false, func(r *sqlx.Rows) error { + for r.Next() { + row := make(map[string]interface{}) + if err := r.MapScan(row); err != nil { + return fmt.Errorf("failed structscan: %w", err) + } + rowsSlice = append(rowsSlice, row) + } + + return nil + }); err != nil { + return fmt.Errorf("querying %s: %w", table, err) + } + + return nil + }); err != nil { + return fmt.Errorf("failed fetch transaction: %w", err) + } + + if len(rowsSlice) == 0 { + break + } + + // Insert + i := sqlite_dialect.Insert(table).Rows(rowsSlice) + sql, args, err := i.ToSQL() + if err != nil { + return fmt.Errorf("failed tosql: %w", err) + } + + _, _, err = db.ExecSQL(ctx, sql, args) + if err != nil { + return fmt.Errorf("exec `%s` [%v]: %w", sql, args, err) + } + + // Move to the next batch + offset += batchSize + } + } + + if err := db.Commit(ctx); err != nil { + return fmt.Errorf("commit: %w", err) + } + + return nil +} diff --git a/pkg/postgres/batch.go b/pkg/postgres/batch.go new file mode 100644 index 0000000000..67f54bb4d0 --- /dev/null +++ b/pkg/postgres/batch.go @@ -0,0 +1,11 @@ +package postgres + +import ( + "github.com/stashapp/stash/pkg/database" +) + +const defaultBatchSize = database.DefaultBatchSize + +func batchExec[T any](ids []T, batchSize int, fn func(batch []T) error) error { + return database.BatchExec(ids, batchSize, fn) +} diff --git a/pkg/postgres/blob.go b/pkg/postgres/blob.go new file mode 100644 index 0000000000..16b3527a64 --- /dev/null +++ b/pkg/postgres/blob.go @@ -0,0 +1,464 @@ +package postgres + +import ( + "context" + "database/sql" + "errors" + "fmt" + "io/fs" + + "github.com/doug-martin/goqu/v9" + "github.com/doug-martin/goqu/v9/exp" + "github.com/jmoiron/sqlx" + "github.com/stashapp/stash/pkg/database" + "github.com/stashapp/stash/pkg/file" + "github.com/stashapp/stash/pkg/hash/md5" + "github.com/stashapp/stash/pkg/logger" + "github.com/stashapp/stash/pkg/sqlite/blob" + "gopkg.in/guregu/null.v4" +) + +const ( + blobTable = "blobs" + blobChecksumColumn = "checksum" +) + +type BlobStore struct { + repository + + tableMgr *table + + fsStore *blob.FilesystemStore + // supplementary stores + otherStores []blob.FilesystemReader + options database.BlobStoreOptions +} + +func NewBlobStore(options database.BlobStoreOptions) *BlobStore { + fs := &file.OsFS{} + + ret := &BlobStore{ + repository: repository{ + tableName: blobTable, + idColumn: blobChecksumColumn, + }, + + tableMgr: blobTableMgr, + + fsStore: blob.NewFilesystemStore(options.Path, fs), + options: options, + } + + for _, otherPath := range options.SupplementaryPaths { + ret.otherStores = append(ret.otherStores, *blob.NewReadonlyFilesystemStore(otherPath, fs)) + } + + return ret +} + +type blobRow struct { + Checksum string `db:"checksum"` + Blob sql.Null[[]byte] `db:"blob"` +} + +func (qb *BlobStore) table() exp.IdentifierExpression { + return qb.tableMgr.table +} + +func (qb *BlobStore) Count(ctx context.Context) (int, error) { + table := qb.table() + q := dialect.From(table).Select(goqu.COUNT(table.Col(blobChecksumColumn))) + + var ret int + if err := querySimple(ctx, q, &ret); err != nil { + return 0, err + } + + return ret, nil +} + +// Write stores the data and its checksum in enabled stores. +// Always writes at least the checksum to the database. +func (qb *BlobStore) Write(ctx context.Context, data []byte) (string, error) { + if !qb.options.UseDatabase && !qb.options.UseFilesystem { + panic("no blob store configured") + } + + if len(data) == 0 { + return "", fmt.Errorf("cannot write empty data") + } + + checksum := md5.FromBytes(data) + + // only write blob to the database if UseDatabase is true + // always at least write the checksum + var storedData sql.Null[[]byte] + if qb.options.UseDatabase { + storedData.V = data + storedData.Valid = len(storedData.V) > 0 + } + + if err := qb.write(ctx, checksum, storedData); err != nil { + return "", fmt.Errorf("writing to database: %w", err) + } + + if qb.options.UseFilesystem { + if err := qb.fsStore.Write(ctx, checksum, data); err != nil { + return "", fmt.Errorf("writing to filesystem: %w", err) + } + } + + return checksum, nil +} + +func (qb *BlobStore) write(ctx context.Context, checksum string, data sql.Null[[]byte]) error { + table := qb.table() + q := dialect.Insert(table).Prepared(true).Rows(blobRow{ + Checksum: checksum, + Blob: data, + }).OnConflict(goqu.DoNothing()) + + _, err := exec(ctx, q) + if err != nil { + return fmt.Errorf("inserting into %s: %w", table, err) + } + + return nil +} + +func (qb *BlobStore) update(ctx context.Context, checksum string, data []byte) error { + table := qb.table() + q := dialect.Update(table).Prepared(true).Set(goqu.Record{ + "blob": data, + }).Where(goqu.C(blobChecksumColumn).Eq(checksum)) + + _, err := exec(ctx, q) + if err != nil { + return fmt.Errorf("updating %s: %w", table, err) + } + + return nil +} + +type ChecksumNotFoundError struct { + Checksum string +} + +func (e *ChecksumNotFoundError) Error() string { + return fmt.Sprintf("checksum %s does not exist", e.Checksum) +} + +type ChecksumBlobNotExistError struct { + Checksum string +} + +func (e *ChecksumBlobNotExistError) Error() string { + return fmt.Sprintf("blob for checksum %s does not exist", e.Checksum) +} + +func (qb *BlobStore) readSQL(ctx context.Context, querySQL sqler) ([]byte, string, error) { + if !qb.options.UseDatabase && !qb.options.UseFilesystem { + panic("no blob store configured") + } + + query, args, err := querySQL.ToSQL() + if err != nil { + return nil, "", fmt.Errorf("reading blob tosql: %w", err) + } + + // always try to get from the database first, even if set to use filesystem + var row blobRow + found := false + const single = true + if err := qb.queryFunc(ctx, query, args, single, func(r *sqlx.Rows) error { + found = true + if err := r.StructScan(&row); err != nil { + return err + } + + return nil + }); err != nil { + return nil, "", fmt.Errorf("reading from database: %w", err) + } + + if !found { + // not found in the database - does not exist + return nil, "", nil + } + + checksum := row.Checksum + + if row.Blob.Valid { + return row.Blob.V, checksum, nil + } + + // don't use the filesystem if not configured to do so + if qb.options.UseFilesystem { + ret, err := qb.readFromFilesystem(ctx, checksum) + if err != nil { + return nil, checksum, err + } + + return ret, checksum, nil + } + + return nil, checksum, &ChecksumBlobNotExistError{ + Checksum: checksum, + } +} + +func (qb *BlobStore) readFromFilesystem(ctx context.Context, checksum string) ([]byte, error) { + // try to read from primary store first, then supplementaries + fsStores := append([]blob.FilesystemReader{qb.fsStore.FilesystemReader}, qb.otherStores...) + + for _, fsStore := range fsStores { + ret, err := fsStore.Read(ctx, checksum) + if err == nil { + return ret, nil + } + + if !errors.Is(err, fs.ErrNotExist) { + return nil, fmt.Errorf("reading from filesystem: %w", err) + } + } + + // blob not found - should not happen + return nil, &ChecksumBlobNotExistError{ + Checksum: checksum, + } +} + +func (qb *BlobStore) EntryExists(ctx context.Context, checksum string) (bool, error) { + q := dialect.From(qb.table()).Select(goqu.COUNT("*")).Where(qb.tableMgr.byID(checksum)) + + var found int + if err := querySimple(ctx, q, &found); err != nil { + return false, fmt.Errorf("querying %s: %w", qb.table(), err) + } + + return found != 0, nil +} + +// Read reads the data from the database or filesystem, depending on which is enabled. +func (qb *BlobStore) Read(ctx context.Context, checksum string) ([]byte, error) { + if !qb.options.UseDatabase && !qb.options.UseFilesystem { + panic("no blob store configured") + } + + // always try to get from the database first, even if set to use filesystem + ret, err := qb.readFromDatabase(ctx, checksum) + if err != nil { + if !errors.Is(err, sql.ErrNoRows) { + return nil, fmt.Errorf("reading from database: %w", err) + } + + // not found in the database - does not exist + return nil, &ChecksumNotFoundError{ + Checksum: checksum, + } + } + + if ret.Valid { + return ret.V, nil + } + + // don't use the filesystem if not configured to do so + if qb.options.UseFilesystem { + return qb.readFromFilesystem(ctx, checksum) + } + + // blob not found - should not happen + return nil, &ChecksumBlobNotExistError{ + Checksum: checksum, + } +} + +func (qb *BlobStore) readFromDatabase(ctx context.Context, checksum string) (sql.Null[[]byte], error) { + q := dialect.From(qb.table()).Select(qb.table().All()).Where(qb.tableMgr.byID(checksum)) + + var row blobRow + const single = true + if err := queryFunc(ctx, q, single, func(r *sqlx.Rows) error { + if err := r.StructScan(&row); err != nil { + return err + } + + return nil + }); err != nil { + return sql.Null[[]byte]{}, fmt.Errorf("querying %s: %w", qb.table(), err) + } + + return row.Blob, nil +} + +// Delete marks a checksum as no longer in use by a single reference. +// If no references remain, the blob is deleted from the database and filesystem. +func (qb *BlobStore) Delete(ctx context.Context, checksum string) error { + // try to delete the blob from the database + if err := qb.delete(ctx, checksum); err != nil { + if qb.isConstraintError(err) { + // blob is still referenced - do not delete + logger.Debugf("Blob %s is still referenced - not deleting", checksum) + return nil + } + + // unexpected error + return fmt.Errorf("deleting from database: %w", err) + } + + // blob was deleted from the database - delete from filesystem if enabled + if qb.options.UseFilesystem { + logger.Debugf("Deleting blob %s from filesystem", checksum) + if err := qb.fsStore.Delete(ctx, checksum); err != nil { + return fmt.Errorf("deleting from filesystem: %w", err) + } + } + + return nil +} + +func (qb *BlobStore) delete(ctx context.Context, checksum string) error { + table := qb.table() + + q := dialect.Delete(table).Where(goqu.C(blobChecksumColumn).Eq(checksum)) + + err := withSavepoint(ctx, func(ctx context.Context) error { + _, err := exec(ctx, q) + return err + }) + if err != nil { + return fmt.Errorf("deleting from %s: %w", table, err) + } + + return nil +} + +type blobJoinQueryBuilder struct { + repository repository + blobStore *BlobStore + + joinTable string +} + +func (qb *blobJoinQueryBuilder) GetImage(ctx context.Context, id int, blobCol string) ([]byte, error) { + sqlQuery := dialect.From(qb.joinTable). + Join(goqu.I("blobs"), goqu.On(goqu.I(qb.joinTable+"."+blobCol).Eq(goqu.I("blobs.checksum")))). + Select(goqu.I("blobs.checksum"), goqu.I("blobs.blob")). + Where(goqu.Ex{"id": id}) + + ret, _, err := qb.blobStore.readSQL(ctx, sqlQuery) + return ret, err +} + +func (qb *blobJoinQueryBuilder) UpdateImage(ctx context.Context, id int, blobCol string, image []byte) error { + if len(image) == 0 { + return qb.DestroyImage(ctx, id, blobCol) + } + + oldChecksum, err := qb.getChecksum(ctx, id, blobCol) + if err != nil { + return err + } + + checksum, err := qb.blobStore.Write(ctx, image) + if err != nil { + return err + } + + sqlQuery := dialect.From(qb.joinTable).Update(). + Set(goqu.Record{blobCol: checksum}). + Prepared(true). + Where(goqu.Ex{"id": id}) + + query, args, err := sqlQuery.ToSQL() + if err != nil { + return err + } + + if _, err := dbWrapper.Exec(ctx, query, args...); err != nil { + return err + } + + // #3595 - delete the old blob if the checksum is different + if oldChecksum != nil && *oldChecksum != checksum { + if err := qb.blobStore.Delete(ctx, *oldChecksum); err != nil { + return err + } + } + + return nil +} + +func (qb *blobJoinQueryBuilder) getChecksum(ctx context.Context, id int, blobCol string) (*string, error) { + sqlQuery := dialect.From(qb.joinTable). + Select(blobCol). + Where(goqu.Ex{"id": id}) + + query, args, err := sqlQuery.ToSQL() + if err != nil { + return nil, err + } + + var checksum null.String + err = qb.repository.querySimple(ctx, query, args, &checksum) + if err != nil { + return nil, err + } + + if !checksum.Valid { + return nil, nil + } + + return &checksum.String, nil +} + +func (qb *blobJoinQueryBuilder) DestroyImage(ctx context.Context, id int, blobCol string) error { + checksum, err := qb.getChecksum(ctx, id, blobCol) + if err != nil { + return err + } + + if checksum == nil { + // no image to delete + return nil + } + + updateQuery := dialect.Update(qb.joinTable). + Set(goqu.Record{blobCol: nil}). + Where(goqu.Ex{"id": id}) + + query, args, err := updateQuery.ToSQL() + if err != nil { + return err + } + + if _, err = dbWrapper.Exec(ctx, query, args...); err != nil { + return err + } + + return qb.blobStore.Delete(ctx, *checksum) +} + +func (qb *blobJoinQueryBuilder) HasImage(ctx context.Context, id int, blobCol string) (bool, error) { + ds := dialect.From(goqu.T(qb.joinTable)). + Select(goqu.C(blobCol)). + Where( + goqu.C("id").Eq(id), + goqu.C(blobCol).IsNotNull(), + ). + Limit(1) + + countDs := dialect.From(ds.As("subquery")).Select(goqu.COUNT("*").As("count")) + + sql, params, err := countDs.ToSQL() + if err != nil { + return false, err + } + + c, err := qb.repository.runCountQuery(ctx, sql, params) + if err != nil { + return false, err + } + + return c == 1, nil +} diff --git a/pkg/postgres/blob_migrate.go b/pkg/postgres/blob_migrate.go new file mode 100644 index 0000000000..1303383772 --- /dev/null +++ b/pkg/postgres/blob_migrate.go @@ -0,0 +1,116 @@ +package postgres + +import ( + "context" + "fmt" + + "github.com/jmoiron/sqlx" +) + +func (qb *BlobStore) FindBlobs(ctx context.Context, n uint, lastChecksum string) ([]string, error) { + table := qb.table() + q := dialect.From(table).Select(table.Col(blobChecksumColumn)).Order(table.Col(blobChecksumColumn).Asc()).Limit(n) + + if lastChecksum != "" { + q = q.Where(table.Col(blobChecksumColumn).Gt(lastChecksum)) + } + + const single = false + var checksums []string + if err := queryFunc(ctx, q, single, func(rows *sqlx.Rows) error { + var checksum string + if err := rows.Scan(&checksum); err != nil { + return err + } + checksums = append(checksums, checksum) + return nil + }); err != nil { + return nil, err + } + + return checksums, nil +} + +// MigrateBlob migrates a blob from the filesystem to the database, or vice versa. +// The target is determined by the UseDatabase and UseFilesystem options. +// If deleteOld is true, the blob is deleted from the source after migration. +func (qb *BlobStore) MigrateBlob(ctx context.Context, checksum string, deleteOld bool) error { + if !qb.options.UseDatabase && !qb.options.UseFilesystem { + panic("no blob store configured") + } + + if qb.options.UseDatabase && qb.options.UseFilesystem { + panic("both filesystem and database configured") + } + + if qb.options.Path == "" { + panic("no blob path configured") + } + + if qb.options.UseDatabase { + return qb.migrateBlobDatabase(ctx, checksum, deleteOld) + } + + return qb.migrateBlobFilesystem(ctx, checksum, deleteOld) +} + +// migrateBlobDatabase migrates a blob from the filesystem to the database +func (qb *BlobStore) migrateBlobDatabase(ctx context.Context, checksum string, deleteOld bool) error { + // ignore if the blob is already present in the database + // (still delete the old data if requested) + existing, err := qb.readFromDatabase(ctx, checksum) + if err != nil { + return fmt.Errorf("reading from database: %w", err) + } + + if len(existing.V) == 0 { + // find the blob in the filesystem + blob, err := qb.fsStore.Read(ctx, checksum) + if err != nil { + return fmt.Errorf("reading from filesystem: %w", err) + } + + // write the blob to the database + if err := qb.update(ctx, checksum, blob); err != nil { + return fmt.Errorf("writing to database: %w", err) + } + } + + if deleteOld { + // delete the blob from the filesystem after commit + if err := qb.fsStore.Delete(ctx, checksum); err != nil { + return fmt.Errorf("deleting from filesystem: %w", err) + } + } + + return nil +} + +// migrateBlobFilesystem migrates a blob from the database to the filesystem +func (qb *BlobStore) migrateBlobFilesystem(ctx context.Context, checksum string, deleteOld bool) error { + // find the blob in the database + blob, err := qb.readFromDatabase(ctx, checksum) + if err != nil { + return fmt.Errorf("reading from database: %w", err) + } + + if len(blob.V) == 0 { + // it's possible that the blob is already present in the filesystem + // just ignore + return nil + } + + // write the blob to the filesystem + if err := qb.fsStore.Write(ctx, checksum, blob.V); err != nil { + return fmt.Errorf("writing to filesystem: %w", err) + } + + if deleteOld { + // delete the blob from the database row + if err := qb.update(ctx, checksum, nil); err != nil { + return err + } + } + + return nil +} diff --git a/pkg/postgres/common.go b/pkg/postgres/common.go new file mode 100644 index 0000000000..f6d5a8b181 --- /dev/null +++ b/pkg/postgres/common.go @@ -0,0 +1,75 @@ +package postgres + +import ( + "context" + "fmt" + + "github.com/doug-martin/goqu/v9" + "github.com/jmoiron/sqlx" +) + +type oCounterManager struct { + tableMgr *table +} + +func (qb *oCounterManager) getOCounter(ctx context.Context, id int) (int, error) { + q := dialect.From(qb.tableMgr.table).Select("o_counter").Where(goqu.Ex{"id": id}) + + const single = true + var ret int + if err := queryFunc(ctx, q, single, func(rows *sqlx.Rows) error { + if err := rows.Scan(&ret); err != nil { + return err + } + return nil + }); err != nil { + return 0, err + } + + return ret, nil +} + +func (qb *oCounterManager) IncrementOCounter(ctx context.Context, id int) (int, error) { + if err := qb.tableMgr.checkIDExists(ctx, id); err != nil { + return 0, err + } + + if err := qb.tableMgr.updateByID(ctx, id, goqu.Record{ + "o_counter": goqu.L("o_counter + 1"), + }); err != nil { + return 0, err + } + + return qb.getOCounter(ctx, id) +} + +func (qb *oCounterManager) DecrementOCounter(ctx context.Context, id int) (int, error) { + if err := qb.tableMgr.checkIDExists(ctx, id); err != nil { + return 0, err + } + + table := qb.tableMgr.table + q := dialect.Update(table).Set(goqu.Record{ + "o_counter": goqu.L("o_counter - 1"), + }).Where(qb.tableMgr.byID(id), goqu.L("o_counter > 0")) + + if _, err := exec(ctx, q); err != nil { + return 0, fmt.Errorf("updating %s: %w", table.GetTable(), err) + } + + return qb.getOCounter(ctx, id) +} + +func (qb *oCounterManager) ResetOCounter(ctx context.Context, id int) (int, error) { + if err := qb.tableMgr.checkIDExists(ctx, id); err != nil { + return 0, err + } + + if err := qb.tableMgr.updateByID(ctx, id, goqu.Record{ + "o_counter": 0, + }); err != nil { + return 0, err + } + + return qb.getOCounter(ctx, id) +} diff --git a/pkg/postgres/criterion_handlers.go b/pkg/postgres/criterion_handlers.go new file mode 100644 index 0000000000..5da68fbbc3 --- /dev/null +++ b/pkg/postgres/criterion_handlers.go @@ -0,0 +1,1170 @@ +package postgres + +import ( + "context" + "database/sql" + "fmt" + "path/filepath" + "regexp" + "strconv" + "strings" + + "github.com/stashapp/stash/pkg/models" + "github.com/stashapp/stash/pkg/utils" +) + +type criterionHandler interface { + handle(ctx context.Context, f *filterBuilder) +} + +type criterionHandlerFunc func(ctx context.Context, f *filterBuilder) + +func (h criterionHandlerFunc) handle(ctx context.Context, f *filterBuilder) { + h(ctx, f) +} + +type compoundHandler []criterionHandler + +func (h compoundHandler) handle(ctx context.Context, f *filterBuilder) { + for _, h := range h { + h.handle(ctx, f) + } +} + +// shared criterion handlers go here + +func stringCriterionHandler(c *models.StringCriterionInput, column string) criterionHandlerFunc { + return func(ctx context.Context, f *filterBuilder) { + if c != nil { + if modifier := c.Modifier; c.Modifier.IsValid() { + switch modifier { + case models.CriterionModifierIncludes: + f.whereClauses = append(f.whereClauses, getStringSearchClause([]string{column}, c.Value, false)) + case models.CriterionModifierExcludes: + f.whereClauses = append(f.whereClauses, getStringSearchClause([]string{column}, c.Value, true)) + case models.CriterionModifierEquals: + f.addWhere(column+" ILIKE ?", c.Value) + case models.CriterionModifierNotEquals: + f.addWhere(column+" NOT ILIKE ?", c.Value) + case models.CriterionModifierMatchesRegex: + if _, err := regexp.Compile(c.Value); err != nil { + f.setError(err) + return + } + f.addWhere(fmt.Sprintf("(%s IS NOT NULL AND regex_match(%[1]s, ?))", column), c.Value) + case models.CriterionModifierNotMatchesRegex: + if _, err := regexp.Compile(c.Value); err != nil { + f.setError(err) + return + } + f.addWhere(fmt.Sprintf("(%s IS NULL OR NOT regex_match(%[1]s, ?))", column), c.Value) + case models.CriterionModifierIsNull: + f.addWhere("(" + column + " IS NULL OR TRIM(" + column + ") = '')") + case models.CriterionModifierNotNull: + f.addWhere("(" + column + " IS NOT NULL AND TRIM(" + column + ") != '')") + default: + panic("unsupported string filter modifier") + } + } + } + } +} + +func uuidCriterionHandler(c *models.StringCriterionInput, column string) criterionHandlerFunc { + return func(ctx context.Context, f *filterBuilder) { + columnCast := "CAST(" + column + " AS TEXT)" + + if c != nil { + if modifier := c.Modifier; c.Modifier.IsValid() { + switch modifier { + case models.CriterionModifierIncludes: + f.whereClauses = append(f.whereClauses, getStringSearchClause([]string{columnCast}, c.Value, false)) + case models.CriterionModifierExcludes: + f.whereClauses = append(f.whereClauses, getStringSearchClause([]string{columnCast}, c.Value, true)) + case models.CriterionModifierEquals: + f.addWhere(columnCast+" ILIKE ?", c.Value) + case models.CriterionModifierNotEquals: + f.addWhere(columnCast+" NOT ILIKE ?", c.Value) + case models.CriterionModifierMatchesRegex: + if _, err := regexp.Compile(c.Value); err != nil { + f.setError(err) + return + } + f.addWhere(fmt.Sprintf("(%s IS NOT NULL AND regex_match(%s, ?))", column, columnCast), c.Value) + case models.CriterionModifierNotMatchesRegex: + if _, err := regexp.Compile(c.Value); err != nil { + f.setError(err) + return + } + f.addWhere(fmt.Sprintf("(%s IS NULL OR NOT regex_match(%s, ?))", column, columnCast), c.Value) + case models.CriterionModifierIsNull: + f.addWhere("(" + column + " IS NULL)") + case models.CriterionModifierNotNull: + f.addWhere("(" + column + " IS NOT NULL)") + default: + panic("unsupported string filter modifier") + } + } + } + } +} + +func joinedStringCriterionHandler(c *models.StringCriterionInput, column string, addJoinFn func(f *filterBuilder)) criterionHandlerFunc { + return func(ctx context.Context, f *filterBuilder) { + if c != nil { + if addJoinFn != nil { + addJoinFn(f) + } + stringCriterionHandler(c, column)(ctx, f) + } + } +} + +func enumCriterionHandler(modifier models.CriterionModifier, values []string, column string) criterionHandlerFunc { + return func(ctx context.Context, f *filterBuilder) { + if modifier.IsValid() { + switch modifier { + case models.CriterionModifierIncludes, models.CriterionModifierEquals: + if len(values) > 0 { + f.whereClauses = append(f.whereClauses, getEnumSearchClause(column, values, false)) + } + case models.CriterionModifierExcludes, models.CriterionModifierNotEquals: + if len(values) > 0 { + f.whereClauses = append(f.whereClauses, getEnumSearchClause(column, values, true)) + } + case models.CriterionModifierIsNull: + f.addWhere("(" + column + " IS NULL OR TRIM(" + column + ") = '')") + case models.CriterionModifierNotNull: + f.addWhere("(" + column + " IS NOT NULL AND TRIM(" + column + ") != '')") + default: + panic("unsupported string filter modifier") + } + } + } +} + +func pathCriterionHandler(c *models.StringCriterionInput, pathColumn string, basenameColumn string, addJoinFn func(f *filterBuilder)) criterionHandlerFunc { + return func(ctx context.Context, f *filterBuilder) { + if c != nil { + if addJoinFn != nil { + addJoinFn(f) + } + addWildcards := true + not := false + + if modifier := c.Modifier; c.Modifier.IsValid() { + switch modifier { + case models.CriterionModifierIncludes: + f.whereClauses = append(f.whereClauses, getPathSearchClauseMany(pathColumn, basenameColumn, c.Value, addWildcards, not)) + case models.CriterionModifierExcludes: + not = true + f.whereClauses = append(f.whereClauses, getPathSearchClauseMany(pathColumn, basenameColumn, c.Value, addWildcards, not)) + case models.CriterionModifierEquals: + addWildcards = false + f.whereClauses = append(f.whereClauses, getPathSearchClause(pathColumn, basenameColumn, c.Value, addWildcards, not)) + case models.CriterionModifierNotEquals: + addWildcards = false + not = true + f.whereClauses = append(f.whereClauses, getPathSearchClause(pathColumn, basenameColumn, c.Value, addWildcards, not)) + case models.CriterionModifierMatchesRegex: + if _, err := regexp.Compile(c.Value); err != nil { + f.setError(err) + return + } + filepathColumn := fmt.Sprintf("%s || '%s' || %s", pathColumn, string(filepath.Separator), basenameColumn) + f.addWhere(fmt.Sprintf("%s IS NOT NULL AND %s IS NOT NULL AND regex_match(%s, ?)", pathColumn, basenameColumn, filepathColumn), c.Value) + case models.CriterionModifierNotMatchesRegex: + if _, err := regexp.Compile(c.Value); err != nil { + f.setError(err) + return + } + filepathColumn := fmt.Sprintf("%s || '%s' || %s", pathColumn, string(filepath.Separator), basenameColumn) + f.addWhere(fmt.Sprintf("%s IS NULL OR %s IS NULL OR NOT regex_match(%s, ?)", pathColumn, basenameColumn, filepathColumn), c.Value) + case models.CriterionModifierIsNull: + f.addWhere(fmt.Sprintf("%s IS NULL OR TRIM(%[1]s) = '' OR %s IS NULL OR TRIM(%[2]s) = ''", pathColumn, basenameColumn)) + case models.CriterionModifierNotNull: + f.addWhere(fmt.Sprintf("%s IS NOT NULL AND TRIM(%[1]s) != '' AND %s IS NOT NULL AND TRIM(%[2]s) != ''", pathColumn, basenameColumn)) + default: + panic("unsupported string filter modifier") + } + } + } + } +} + +func getPathSearchClause(pathColumn, basenameColumn, p string, addWildcards, not bool) sqlClause { + if addWildcards { + p = "%" + p + "%" + } + + filepathColumn := fmt.Sprintf("%s || '%s' || %s", pathColumn, string(filepath.Separator), basenameColumn) + ret := makeClause(fmt.Sprintf("%s ILIKE ?", filepathColumn), p) + + if not { + ret = ret.not() + } + + return ret +} + +// getPathSearchClauseMany splits the query string p on whitespace +// Used for backwards compatibility for the includes/excludes modifiers +func getPathSearchClauseMany(pathColumn, basenameColumn, p string, addWildcards, not bool) sqlClause { + q := strings.TrimSpace(p) + trimmedQuery := strings.Trim(q, "\"") + + if trimmedQuery == q { + q = regexp.MustCompile(`\s+`).ReplaceAllString(q, " ") + queryWords := strings.Split(q, " ") + + var ret []sqlClause + // Search for any word + for _, word := range queryWords { + ret = append(ret, getPathSearchClause(pathColumn, basenameColumn, word, addWildcards, not)) + } + + if !not { + return orClauses(ret...) + } + + return andClauses(ret...) + } + + return getPathSearchClause(pathColumn, basenameColumn, trimmedQuery, addWildcards, not) +} + +func intCriterionHandler(c *models.IntCriterionInput, column string, addJoinFn func(f *filterBuilder)) criterionHandlerFunc { + return func(ctx context.Context, f *filterBuilder) { + if c != nil { + if addJoinFn != nil { + addJoinFn(f) + } + clause, args := getIntCriterionWhereClause(column, *c) + f.addWhere(clause, args...) + } + } +} + +func floatCriterionHandler(c *models.FloatCriterionInput, column string, addJoinFn func(f *filterBuilder)) criterionHandlerFunc { + return func(ctx context.Context, f *filterBuilder) { + if c != nil { + if addJoinFn != nil { + addJoinFn(f) + } + clause, args := getFloatCriterionWhereClause(column, *c) + f.addWhere(clause, args...) + } + } +} + +func floatIntCriterionHandler(durationFilter *models.IntCriterionInput, column string, addJoinFn func(f *filterBuilder)) criterionHandlerFunc { + return func(ctx context.Context, f *filterBuilder) { + if durationFilter != nil { + if addJoinFn != nil { + addJoinFn(f) + } + clause, args := getIntCriterionWhereClause("cast("+column+" as int)", *durationFilter) + f.addWhere(clause, args...) + } + } +} + +func boolCriterionHandler(c *bool, column string, addJoinFn func(f *filterBuilder)) criterionHandlerFunc { + return func(ctx context.Context, f *filterBuilder) { + if c != nil { + if addJoinFn != nil { + addJoinFn(f) + } + var v = strconv.FormatBool(*c) + + f.addWhere(column + " = " + v) + } + } +} + +type dateCriterionHandler struct { + c *models.DateCriterionInput + column string + joinFn func(f *filterBuilder) +} + +func (h *dateCriterionHandler) handle(ctx context.Context, f *filterBuilder) { + if h.c != nil { + if h.joinFn != nil { + h.joinFn(f) + } + clause, args := getDateCriterionWhereClause(h.column, *h.c) + f.addWhere(clause, args...) + } +} + +type timestampCriterionHandler struct { + c *models.TimestampCriterionInput + column string + joinFn func(f *filterBuilder) +} + +func (h *timestampCriterionHandler) handle(ctx context.Context, f *filterBuilder) { + if h.c != nil { + if h.joinFn != nil { + h.joinFn(f) + } + clause, args := getTimestampCriterionWhereClause(h.column, *h.c) + f.addWhere(clause, args...) + } +} + +func yearFilterCriterionHandler(year *models.IntCriterionInput, col string) criterionHandlerFunc { + return func(ctx context.Context, f *filterBuilder) { + if year != nil && year.Modifier.IsValid() { + clause, args := getIntCriterionWhereClause("TO_CHAR("+col+", 'YYYY')::int", *year) + f.addWhere(clause, args...) + } + } +} + +func resolutionCriterionHandler(resolution *models.ResolutionCriterionInput, heightColumn string, widthColumn string, addJoinFn func(f *filterBuilder)) criterionHandlerFunc { + return func(ctx context.Context, f *filterBuilder) { + if resolution != nil && resolution.Value.IsValid() { + if addJoinFn != nil { + addJoinFn(f) + } + + mn := resolution.Value.GetMinResolution() + mx := resolution.Value.GetMaxResolution() + + widthHeight := fmt.Sprintf("LEAST(%s, %s)", widthColumn, heightColumn) + + switch resolution.Modifier { + case models.CriterionModifierEquals: + f.addWhere(fmt.Sprintf("%s BETWEEN %d AND %d", widthHeight, mn, mx)) + case models.CriterionModifierNotEquals: + f.addWhere(fmt.Sprintf("%s NOT BETWEEN %d AND %d", widthHeight, mn, mx)) + case models.CriterionModifierLessThan: + f.addWhere(fmt.Sprintf("%s < %d", widthHeight, mn)) + case models.CriterionModifierGreaterThan: + f.addWhere(fmt.Sprintf("%s > %d", widthHeight, mx)) + } + } + } +} + +func orientationCriterionHandler(orientation *models.OrientationCriterionInput, heightColumn string, widthColumn string, addJoinFn func(f *filterBuilder)) criterionHandlerFunc { + return func(ctx context.Context, f *filterBuilder) { + if orientation != nil { + if addJoinFn != nil { + addJoinFn(f) + } + + var clauses []sqlClause + + for _, v := range orientation.Value { + // width mod height + mod := "" + switch v { + case models.OrientationPortrait: + mod = "<" + case models.OrientationLandscape: + mod = ">" + case models.OrientationSquare: + mod = "=" + } + + if mod != "" { + clauses = append(clauses, makeClause(fmt.Sprintf("%s %s %s", widthColumn, mod, heightColumn))) + } + } + + if len(clauses) > 0 { + f.whereClauses = append(f.whereClauses, orClauses(clauses...)) + } + } + } +} + +// handle for MultiCriterion where there is a join table between the new +// objects +type joinedMultiCriterionHandlerBuilder struct { + // table containing the primary objects + primaryTable string + // table joining primary and foreign objects + joinTable string + // alias for join table, if required + joinAs string + // foreign key of the primary object on the join table + primaryFK string + // foreign key of the foreign object on the join table + foreignFK string + + addJoinTable func(f *filterBuilder) +} + +func (m *joinedMultiCriterionHandlerBuilder) handler(c *models.MultiCriterionInput) criterionHandlerFunc { + return func(ctx context.Context, f *filterBuilder) { + if c != nil { + // make local copy so we can modify it + criterion := *c + + joinAlias := m.joinAs + if joinAlias == "" { + joinAlias = m.joinTable + } + + if criterion.Modifier == models.CriterionModifierIsNull || criterion.Modifier == models.CriterionModifierNotNull { + var notClause string + if criterion.Modifier == models.CriterionModifierNotNull { + notClause = "NOT" + } + + m.addJoinTable(f) + + f.addWhere(utils.StrFormat("{table}.{column} IS {not} NULL", utils.StrFormatMap{ + "table": joinAlias, + "column": m.foreignFK, + "not": notClause, + })) + return + } + + if len(criterion.Value) == 0 && len(criterion.Excludes) == 0 { + return + } + + // combine excludes if excludes modifier is selected + if criterion.Modifier == models.CriterionModifierExcludes { + criterion.Modifier = models.CriterionModifierIncludesAll + criterion.Excludes = append(criterion.Excludes, criterion.Value...) + criterion.Value = nil + } + + if len(criterion.Value) > 0 { + whereClause := "" + havingClause := "" + + var args []interface{} + for _, tagID := range criterion.Value { + args = append(args, tagID) + } + + switch criterion.Modifier { + case models.CriterionModifierIncludes: + // includes any of the provided ids + m.addJoinTable(f) + whereClause = fmt.Sprintf("%s.%s IN %s", joinAlias, m.foreignFK, getInBinding(len(criterion.Value))) + case models.CriterionModifierEquals: + // includes only the provided ids + m.addJoinTable(f) + whereClause = utils.StrFormat("{joinAlias}.{foreignFK} IN {inBinding} AND (SELECT COUNT(*) FROM {joinTable} s WHERE s.{primaryFK} = {primaryTable}.id) = ?", utils.StrFormatMap{ + "joinAlias": joinAlias, + "foreignFK": m.foreignFK, + "inBinding": getInBinding(len(criterion.Value)), + "joinTable": m.joinTable, + "primaryFK": m.primaryFK, + "primaryTable": m.primaryTable, + }) + havingClause = fmt.Sprintf("count(distinct %s.%s) = %d", joinAlias, m.foreignFK, len(criterion.Value)) + args = append(args, len(criterion.Value)) + case models.CriterionModifierNotEquals: + f.setError(fmt.Errorf("not equals modifier is not supported for multi criterion input")) + case models.CriterionModifierIncludesAll: + // includes all of the provided ids + m.addJoinTable(f) + whereClause = fmt.Sprintf("%s.%s IN %s", joinAlias, m.foreignFK, getInBinding(len(criterion.Value))) + havingClause = fmt.Sprintf("count(distinct %s.%s) = %d", joinAlias, m.foreignFK, len(criterion.Value)) + } + + f.addWhere(whereClause, args...) + f.addHaving(havingClause) + } + + if len(criterion.Excludes) > 0 { + var args []interface{} + for _, tagID := range criterion.Excludes { + args = append(args, tagID) + } + + // excludes all of the provided ids + // need to use actual join table name for this + // .id NOT IN (select . from where . in ) + whereClause := fmt.Sprintf("%[1]s.id NOT IN (SELECT %[3]s.%[2]s from %[3]s where %[3]s.%[4]s in %[5]s)", m.primaryTable, m.primaryFK, m.joinTable, m.foreignFK, getInBinding(len(criterion.Excludes))) + + f.addWhere(whereClause, args...) + } + } + } +} + +type multiCriterionHandlerBuilder struct { + primaryTable string + foreignTable string + joinTable string + primaryFK string + foreignFK string + + // function that will be called to perform any necessary joins + addJoinsFunc func(f *filterBuilder) +} + +func (m *multiCriterionHandlerBuilder) handler(criterion *models.MultiCriterionInput) criterionHandlerFunc { + return func(ctx context.Context, f *filterBuilder) { + if criterion != nil { + if criterion.Modifier == models.CriterionModifierIsNull || criterion.Modifier == models.CriterionModifierNotNull { + var notClause string + if criterion.Modifier == models.CriterionModifierNotNull { + notClause = "NOT" + } + + table := m.primaryTable + if m.joinTable != "" { + table = m.joinTable + f.addLeftJoin(table, "", fmt.Sprintf("%s.%s = %s.id", table, m.primaryFK, m.primaryTable)) + } + + f.addWhere(fmt.Sprintf("%s.%s IS %s NULL", table, m.foreignFK, notClause)) + return + } + + if len(criterion.Value) == 0 { + return + } + + var args []interface{} + for _, tagID := range criterion.Value { + args = append(args, tagID) + } + + if m.addJoinsFunc != nil { + m.addJoinsFunc(f) + } + + whereClause, havingClause := getMultiCriterionClause(m.primaryTable, m.foreignTable, m.joinTable, m.primaryFK, m.foreignFK, criterion) + f.addWhere(whereClause, args...) + f.addHaving(havingClause) + } + } +} + +type countCriterionHandlerBuilder struct { + primaryTable string + joinTable string + primaryFK string +} + +func (m *countCriterionHandlerBuilder) handler(criterion *models.IntCriterionInput) criterionHandlerFunc { + return func(ctx context.Context, f *filterBuilder) { + if criterion != nil { + clause, args := getCountCriterionClause(m.primaryTable, m.joinTable, m.primaryFK, *criterion) + + f.addWhere(clause, args...) + } + } +} + +// handler for StringCriterion for string list fields +type stringListCriterionHandlerBuilder struct { + primaryTable string + // foreign key of the primary object on the join table + primaryFK string + // table joining primary and foreign objects + joinTable string + // string field on the join table + stringColumn string + + addJoinTable func(f *filterBuilder) + excludeHandler func(f *filterBuilder, criterion *models.StringCriterionInput) +} + +func (m *stringListCriterionHandlerBuilder) handler(criterion *models.StringCriterionInput) criterionHandlerFunc { + return func(ctx context.Context, f *filterBuilder) { + if criterion != nil { + if criterion.Modifier == models.CriterionModifierExcludes { + // special handling for excludes + if m.excludeHandler != nil { + m.excludeHandler(f, criterion) + return + } + + // excludes all of the provided values + // need to use actual join table name for this + // .id NOT IN (select . from where . in ) + whereClause := utils.StrFormat("{primaryTable}.id NOT IN (SELECT {joinTable}.{primaryFK} from {joinTable} where {joinTable}.{stringColumn} ILIKE ?)", + utils.StrFormatMap{ + "primaryTable": m.primaryTable, + "joinTable": m.joinTable, + "primaryFK": m.primaryFK, + "stringColumn": m.stringColumn, + }, + ) + + f.addWhere(whereClause, "%"+criterion.Value+"%") + + // TODO - should we also exclude null values? + // m.addJoinTable(f) + // stringCriterionHandler(&models.StringCriterionInput{ + // Modifier: models.CriterionModifierNotNull, + // }, m.joinTable+"."+m.stringColumn)(ctx, f) + } else { + m.addJoinTable(f) + stringCriterionHandler(criterion, m.joinTable+"."+m.stringColumn)(ctx, f) + } + } + } +} + +func studioCriterionHandler(primaryTable string, studios *models.HierarchicalMultiCriterionInput) criterionHandlerFunc { + return func(ctx context.Context, f *filterBuilder) { + if studios == nil { + return + } + + studiosCopy := *studios + switch studiosCopy.Modifier { + case models.CriterionModifierEquals: + studiosCopy.Modifier = models.CriterionModifierIncludesAll + case models.CriterionModifierNotEquals: + studiosCopy.Modifier = models.CriterionModifierExcludes + } + + hh := hierarchicalMultiCriterionHandlerBuilder{ + primaryTable: primaryTable, + foreignTable: studioTable, + foreignFK: studioIDColumn, + parentFK: "parent_id", + } + + hh.handler(&studiosCopy)(ctx, f) + } +} + +type hierarchicalMultiCriterionHandlerBuilder struct { + primaryTable string + foreignTable string + foreignFK string + + parentFK string + childFK string + relationsTable string +} + +func getHierarchicalValues(ctx context.Context, values []string, table, relationsTable, parentFK string, childFK string, depth *int, parenthesis bool) (string, error) { + var args []interface{} + + if parentFK == "" { + parentFK = "parent_id" + } + if childFK == "" { + childFK = "child_id" + } + + depthVal := 0 + if depth != nil { + depthVal = *depth + } + + if depthVal == 0 { + valid := true + var valuesClauses []string + for _, value := range values { + id, err := strconv.Atoi(value) + // In case of invalid value just run the query. + // Building VALUES() based on provided values just saves a query when depth is 0. + if err != nil { + valid = false + break + } + + valuesClauses = append(valuesClauses, fmt.Sprintf("(%d,%d)", id, id)) + } + + if valid { + values := "VALUES" + strings.Join(valuesClauses, ",") + if parenthesis { + values = "(" + values + ") AS v(column1, column2)" + } + return values, nil + } + } + + for _, value := range values { + args = append(args, value) + } + inCount := len(args) + + var depthCondition string + if depthVal != -1 { + depthCondition = fmt.Sprintf("WHERE depth < %d", depthVal) + } + + withClauseMap := utils.StrFormatMap{ + "table": table, + "relationsTable": relationsTable, + "inBinding": getInBinding(inCount), + "recursiveSelect": "", + "parentFK": parentFK, + "childFK": childFK, + "depthCondition": depthCondition, + "unionClause": "", + } + + if relationsTable != "" { + withClauseMap["recursiveSelect"] = utils.StrFormat(`SELECT p.root_id, c.{childFK}, depth + 1 FROM {relationsTable} AS c +INNER JOIN items as p ON c.{parentFK} = p.item_id +`, withClauseMap) + } else { + withClauseMap["recursiveSelect"] = utils.StrFormat(`SELECT p.root_id, c.id, depth + 1 FROM {table} as c +INNER JOIN items as p ON c.{parentFK} = p.item_id +`, withClauseMap) + } + + if depthVal != 0 { + withClauseMap["unionClause"] = utils.StrFormat(` +UNION {recursiveSelect} {depthCondition} +`, withClauseMap) + } + + withClause := utils.StrFormat(`items AS ( +SELECT id as root_id, id as item_id, 0 as depth FROM {table} +WHERE id in {inBinding} +{unionClause}) +`, withClauseMap) + + query := fmt.Sprintf("WITH RECURSIVE %s SELECT 'VALUES' || STRING_AGG('(' || root_id || ', ' || item_id || ')'::TEXT, ',') AS val FROM items", withClause) + + var valuesClause sql.NullString + err := dbWrapper.Get(ctx, &valuesClause, query, args...) + if err != nil { + return "", fmt.Errorf("failed to get hierarchical values: %w", err) + } + + // if no values are found, just return a values string with the values only + if !valuesClause.Valid { + for i, value := range values { + values[i] = fmt.Sprintf("(%s, %s)", value, value) + } + valuesClause.String = "VALUES" + strings.Join(values, ",") + } + + if parenthesis { + valuesClause.String = "(" + valuesClause.String + ") AS v(column1, column2)" + } + + return valuesClause.String, nil +} + +func addHierarchicalConditionClauses(f *filterBuilder, criterion models.HierarchicalMultiCriterionInput, table, idColumn string) { + switch criterion.Modifier { + case models.CriterionModifierIncludes: + f.addWhere(fmt.Sprintf("%s.%s IS NOT NULL", table, idColumn)) + case models.CriterionModifierIncludesAll: + f.addWhere(fmt.Sprintf("%s.%s IS NOT NULL", table, idColumn)) + f.addHaving(fmt.Sprintf("count(distinct %s.%s) = %d", table, idColumn, len(criterion.Value))) + case models.CriterionModifierExcludes: + f.addWhere(fmt.Sprintf("%s.%s IS NULL", table, idColumn)) + } +} + +func (m *hierarchicalMultiCriterionHandlerBuilder) handler(c *models.HierarchicalMultiCriterionInput) criterionHandlerFunc { + return func(ctx context.Context, f *filterBuilder) { + if c != nil { + // make a copy so we don't modify the original + criterion := *c + + // don't support equals/not equals + if criterion.Modifier == models.CriterionModifierEquals || criterion.Modifier == models.CriterionModifierNotEquals { + f.setError(fmt.Errorf("modifier %s is not supported for hierarchical multi criterion", criterion.Modifier)) + return + } + + if criterion.Modifier == models.CriterionModifierIsNull || criterion.Modifier == models.CriterionModifierNotNull { + var notClause string + if criterion.Modifier == models.CriterionModifierNotNull { + notClause = "NOT" + } + + f.addWhere(utils.StrFormat("{table}.{column} IS {not} NULL", utils.StrFormatMap{ + "table": m.primaryTable, + "column": m.foreignFK, + "not": notClause, + })) + return + } + + if len(criterion.Value) == 0 && len(criterion.Excludes) == 0 { + return + } + + // combine excludes if excludes modifier is selected + if criterion.Modifier == models.CriterionModifierExcludes { + criterion.Modifier = models.CriterionModifierIncludesAll + criterion.Excludes = append(criterion.Excludes, criterion.Value...) + criterion.Value = nil + } + + if len(criterion.Value) > 0 { + valuesClause, err := getHierarchicalValues(ctx, criterion.Value, m.foreignTable, m.relationsTable, m.parentFK, m.childFK, criterion.Depth, true) + if err != nil { + f.setError(err) + return + } + + switch criterion.Modifier { + case models.CriterionModifierIncludes: + f.addWhere(fmt.Sprintf("%s.%s IN (SELECT column2 FROM %s)", m.primaryTable, m.foreignFK, valuesClause)) + case models.CriterionModifierIncludesAll: + f.addWhere(fmt.Sprintf("%s.%s IN (SELECT column2 FROM %s)", m.primaryTable, m.foreignFK, valuesClause)) + f.addHaving(fmt.Sprintf("count(distinct %s.%s) = %d", m.primaryTable, m.foreignFK, len(criterion.Value))) + } + } + + if len(criterion.Excludes) > 0 { + valuesClause, err := getHierarchicalValues(ctx, criterion.Excludes, m.foreignTable, m.relationsTable, m.parentFK, m.childFK, criterion.Depth, true) + if err != nil { + f.setError(err) + return + } + + f.addWhere(fmt.Sprintf("%s.%s NOT IN (SELECT column2 FROM %s) OR %[1]s.%[2]s IS NULL", m.primaryTable, m.foreignFK, valuesClause)) + } + } + } +} + +type joinedHierarchicalMultiCriterionHandlerBuilder struct { + primaryTable string + primaryKey string + foreignTable string + foreignFK string + + parentFK string + childFK string + relationsTable string + + joinAs string + joinTable string + primaryFK string +} + +func (m *joinedHierarchicalMultiCriterionHandlerBuilder) addHierarchicalConditionClauses(f *filterBuilder, criterion models.HierarchicalMultiCriterionInput, table, idColumn string) { + primaryKey := m.primaryKey + if primaryKey == "" { + primaryKey = "id" + } + + switch criterion.Modifier { + case models.CriterionModifierEquals: + // includes only the provided ids + f.addWhere(fmt.Sprintf("%s.%s IS NOT NULL", table, idColumn)) + f.addHaving(fmt.Sprintf("count(distinct %s.%s) = %d", table, idColumn, len(criterion.Value))) + f.addWhere(utils.StrFormat("(SELECT COUNT(*) FROM {joinTable} s WHERE s.{primaryFK} = {primaryTable}.{primaryKey}) = ?", utils.StrFormatMap{ + "joinTable": m.joinTable, + "primaryFK": m.primaryFK, + "primaryTable": m.primaryTable, + "primaryKey": primaryKey, + }), len(criterion.Value)) + case models.CriterionModifierNotEquals: + f.setError(fmt.Errorf("not equals modifier is not supported for hierarchical multi criterion input")) + default: + addHierarchicalConditionClauses(f, criterion, table, idColumn) + } +} + +func (m *joinedHierarchicalMultiCriterionHandlerBuilder) handler(c *models.HierarchicalMultiCriterionInput) criterionHandlerFunc { + return func(ctx context.Context, f *filterBuilder) { + if c != nil { + // make a copy so we don't modify the original + criterion := *c + joinAlias := m.joinAs + primaryKey := m.primaryKey + if primaryKey == "" { + primaryKey = "id" + } + + if criterion.Modifier == models.CriterionModifierEquals && criterion.Depth != nil && *criterion.Depth != 0 { + f.setError(fmt.Errorf("depth is not supported for equals modifier in hierarchical multi criterion input")) + return + } + + if criterion.Modifier == models.CriterionModifierIsNull || criterion.Modifier == models.CriterionModifierNotNull { + var notClause string + if criterion.Modifier == models.CriterionModifierNotNull { + notClause = "NOT" + } + + f.addLeftJoin(m.joinTable, joinAlias, fmt.Sprintf("%s.%s = %s.%s", joinAlias, m.primaryFK, m.primaryTable, primaryKey)) + + f.addWhere(utils.StrFormat("{table}.{column} IS {not} NULL", utils.StrFormatMap{ + "table": joinAlias, + "column": m.foreignFK, + "not": notClause, + })) + return + } + + // combine excludes if excludes modifier is selected + if criterion.Modifier == models.CriterionModifierExcludes { + criterion.Modifier = models.CriterionModifierIncludesAll + criterion.Excludes = append(criterion.Excludes, criterion.Value...) + criterion.Value = nil + } + + if len(criterion.Value) == 0 && len(criterion.Excludes) == 0 { + return + } + + if len(criterion.Value) > 0 { + valuesClause, err := getHierarchicalValues(ctx, criterion.Value, m.foreignTable, m.relationsTable, m.parentFK, m.childFK, criterion.Depth, false) + if err != nil { + f.setError(err) + return + } + + joinTable := utils.StrFormat(`( + SELECT j.*, d.column1 AS root_id, d.column2 AS item_id FROM {joinTable} AS j + INNER JOIN ({valuesClause}) AS d ON j.{foreignFK} = d.column2 + ) + `, utils.StrFormatMap{ + "joinTable": m.joinTable, + "foreignFK": m.foreignFK, + "valuesClause": valuesClause, + }) + + f.addLeftJoin(joinTable, joinAlias, fmt.Sprintf("%s.%s = %s.%s", joinAlias, m.primaryFK, m.primaryTable, primaryKey)) + + m.addHierarchicalConditionClauses(f, criterion, joinAlias, "root_id") + } + + if len(criterion.Excludes) > 0 { + valuesClause, err := getHierarchicalValues(ctx, criterion.Excludes, m.foreignTable, m.relationsTable, m.parentFK, m.childFK, criterion.Depth, false) + if err != nil { + f.setError(err) + return + } + + joinTable := utils.StrFormat(`( + SELECT j2.*, e.column1 AS root_id, e.column2 AS item_id FROM {joinTable} AS j2 + INNER JOIN ({valuesClause}) AS e ON j2.{foreignFK} = e.column2 + ) + `, utils.StrFormatMap{ + "joinTable": m.joinTable, + "foreignFK": m.foreignFK, + "valuesClause": valuesClause, + }) + + joinAlias2 := joinAlias + "2" + + f.addLeftJoin(joinTable, joinAlias2, fmt.Sprintf("%s.%s = %s.%s", joinAlias2, m.primaryFK, m.primaryTable, primaryKey)) + + // modify for exclusion + criterionCopy := criterion + criterionCopy.Modifier = models.CriterionModifierExcludes + criterionCopy.Value = c.Excludes + + m.addHierarchicalConditionClauses(f, criterionCopy, joinAlias2, "root_id") + } + } + } +} + +type joinedPerformerTagsHandler struct { + criterion *models.HierarchicalMultiCriterionInput + + primaryTable string // eg scenes + joinTable string // eg performers_scenes + joinPrimaryKey string // eg scene_id +} + +func (h *joinedPerformerTagsHandler) handle(ctx context.Context, f *filterBuilder) { + tags := h.criterion + + if tags != nil { + criterion := tags.CombineExcludes() + + // validate the modifier + switch criterion.Modifier { + case models.CriterionModifierIncludesAll, models.CriterionModifierIncludes, models.CriterionModifierExcludes, models.CriterionModifierIsNull, models.CriterionModifierNotNull: + // valid + default: + f.setError(fmt.Errorf("invalid modifier %s for performer tags", criterion.Modifier)) + } + + strFormatMap := utils.StrFormatMap{ + "primaryTable": h.primaryTable, + "joinTable": h.joinTable, + "joinPrimaryKey": h.joinPrimaryKey, + "inBinding": getInBinding(len(criterion.Value)), + } + + if criterion.Modifier == models.CriterionModifierIsNull || criterion.Modifier == models.CriterionModifierNotNull { + var notClause string + if criterion.Modifier == models.CriterionModifierNotNull { + notClause = "NOT" + } + + f.addLeftJoin(h.joinTable, "", utils.StrFormat("{primaryTable}.id = {joinTable}.{joinPrimaryKey}", strFormatMap)) + f.addLeftJoin("performers_tags", "", utils.StrFormat("{joinTable}.performer_id = performers_tags.performer_id", strFormatMap)) + + f.addWhere(fmt.Sprintf("performers_tags.tag_id IS %s NULL", notClause)) + return + } + + if len(criterion.Value) == 0 && len(criterion.Excludes) == 0 { + return + } + + if len(criterion.Value) > 0 { + valuesClause, err := getHierarchicalValues(ctx, criterion.Value, tagTable, "tags_relations", "", "", criterion.Depth, false) + if err != nil { + f.setError(err) + return + } + + f.addWith(utils.StrFormat(`performer_tags AS ( +SELECT ps.{joinPrimaryKey} as primaryID, t.column1 AS root_tag_id FROM {joinTable} ps +INNER JOIN performers_tags pt ON pt.performer_id = ps.performer_id +INNER JOIN (`+valuesClause+`) t ON t.column2 = pt.tag_id +)`, strFormatMap)) + + f.addLeftJoin("performer_tags", "", utils.StrFormat("performer_tags.primaryID = {primaryTable}.id", strFormatMap)) + + addHierarchicalConditionClauses(f, criterion, "performer_tags", "root_tag_id") + } + + if len(criterion.Excludes) > 0 { + valuesClause, err := getHierarchicalValues(ctx, criterion.Excludes, tagTable, "tags_relations", "", "", criterion.Depth, true) + if err != nil { + f.setError(err) + return + } + + clause := utils.StrFormat("{primaryTable}.id NOT IN (SELECT {joinTable}.{joinPrimaryKey} FROM {joinTable} INNER JOIN performers_tags ON {joinTable}.performer_id = performers_tags.performer_id WHERE performers_tags.tag_id IN (SELECT column2 FROM %s))", strFormatMap) + f.addWhere(fmt.Sprintf(clause, valuesClause)) + } + } +} + +type stashIDCriterionHandler struct { + c *models.StashIDCriterionInput + stashIDRepository *stashIDRepository + stashIDTableAs string + parentIDCol string +} + +func (h *stashIDCriterionHandler) handle(ctx context.Context, f *filterBuilder) { + if h.c == nil { + return + } + + // ideally, this handler should just convert to stashIDsCriterionHandler + // but there are some differences in how the existing handler works compared + // to the new code, specifically because this code uses the stringCriterionHandler. + // To minimise potential regressions, we'll keep the existing logic for now. + + stashIDRepo := h.stashIDRepository + t := stashIDRepo.tableName + if h.stashIDTableAs != "" { + t = h.stashIDTableAs + } + + joinClause := fmt.Sprintf("%s.%s = %s", t, stashIDRepo.idColumn, h.parentIDCol) + if h.c.Endpoint != nil && *h.c.Endpoint != "" { + joinClause += fmt.Sprintf(" AND %s.endpoint = '%s'", t, *h.c.Endpoint) + } + + f.addLeftJoin(stashIDRepo.tableName, h.stashIDTableAs, joinClause) + + v := "" + if h.c.StashID != nil { + v = *h.c.StashID + } + + uuidCriterionHandler(&models.StringCriterionInput{ + Value: v, + Modifier: h.c.Modifier, + }, t+".stash_id")(ctx, f) +} + +type stashIDsCriterionHandler struct { + c *models.StashIDsCriterionInput + stashIDRepository *stashIDRepository + stashIDTableAs string + parentIDCol string +} + +func (h *stashIDsCriterionHandler) handle(ctx context.Context, f *filterBuilder) { + if h.c == nil { + return + } + + stashIDRepo := h.stashIDRepository + t := stashIDRepo.tableName + if h.stashIDTableAs != "" { + t = h.stashIDTableAs + } + + joinClause := fmt.Sprintf("%s.%s = %s", t, stashIDRepo.idColumn, h.parentIDCol) + if h.c.Endpoint != nil && *h.c.Endpoint != "" { + joinClause += fmt.Sprintf(" AND %s.endpoint = '%s'", t, *h.c.Endpoint) + } + + f.addLeftJoin(stashIDRepo.tableName, h.stashIDTableAs, joinClause) + + switch h.c.Modifier { + case models.CriterionModifierIsNull: + f.addWhere(fmt.Sprintf("%s.stash_id IS NULL", t)) + case models.CriterionModifierNotNull: + f.addWhere(fmt.Sprintf("%s.stash_id IS NOT NULL", t)) + case models.CriterionModifierEquals: + var clauses []sqlClause + for _, id := range h.c.StashIDs { + clauses = append(clauses, makeClause(fmt.Sprintf("%s.stash_id = '?'::uuid", t), id)) + } + f.whereClauses = append(f.whereClauses, orClauses(clauses...)) + case models.CriterionModifierNotEquals: + var clauses []sqlClause + for _, id := range h.c.StashIDs { + clauses = append(clauses, makeClause(fmt.Sprintf("%s.stash_id != '?'::uuid", t), id)) + } + f.whereClauses = append(f.whereClauses, andClauses(clauses...)) + default: + f.setError(fmt.Errorf("invalid modifier %s for stash IDs criterion", h.c.Modifier)) + } +} + +type relatedFilterHandler struct { + relatedIDCol string + relatedRepo repository + relatedHandler criterionHandler + joinFn func(f *filterBuilder) + directJoin bool +} + +func (h *relatedFilterHandler) handle(ctx context.Context, f *filterBuilder) { + ff := filterBuilderFromHandler(ctx, h.relatedHandler) + if ff.err != nil { + f.setError(ff.err) + return + } + + if ff.empty() { + return + } + + if h.joinFn != nil { + h.joinFn(f) + } + + if h.directJoin { + // rerun handler using existing filter builder + h.relatedHandler.handle(ctx, f) + return + } + + subQuery := h.relatedRepo.newQuery() + selectIDs(&subQuery, subQuery.repository.tableName) + if err := subQuery.addFilter(ff); err != nil { + f.setError(err) + return + } + + f.addWhere(fmt.Sprintf("%s IN ("+subQuery.toSQL(false)+")", h.relatedIDCol), subQuery.args...) +} diff --git a/pkg/postgres/custom_fields.go b/pkg/postgres/custom_fields.go new file mode 100644 index 0000000000..5c5510ba07 --- /dev/null +++ b/pkg/postgres/custom_fields.go @@ -0,0 +1,398 @@ +package postgres + +import ( + "context" + "encoding/json" + "fmt" + "reflect" + "regexp" + "strings" + + "github.com/doug-martin/goqu/v9" + "github.com/doug-martin/goqu/v9/exp" + "github.com/jmoiron/sqlx" + "github.com/stashapp/stash/pkg/models" +) + +const maxCustomFieldNameLength = 64 + +type customFieldsStore struct { + table exp.IdentifierExpression + fk exp.IdentifierExpression +} + +func (s *customFieldsStore) deleteForID(ctx context.Context, id int) error { + table := s.table + q := dialect.Delete(table).Where(s.fk.Eq(id)) + _, err := exec(ctx, q) + if err != nil { + return fmt.Errorf("deleting from %s: %w", s.table.GetTable(), err) + } + + return nil +} + +func (s *customFieldsStore) SetCustomFields(ctx context.Context, id int, values models.CustomFieldsInput) error { + var partial bool + var valMap map[string]interface{} + + switch { + case values.Full != nil: + partial = false + valMap = values.Full + case values.Partial != nil: + partial = true + valMap = values.Partial + } + + if valMap != nil { + if err := s.validateCustomFields(valMap, values.Remove); err != nil { + return err + } + + if err := s.setCustomFields(ctx, id, valMap, partial); err != nil { + return err + } + } + + if err := s.deleteCustomFields(ctx, id, values.Remove); err != nil { + return err + } + + return nil +} + +func (s *customFieldsStore) validateCustomFields(values map[string]interface{}, deleteKeys []string) error { + // if values is nil, nothing to validate + if values == nil { + return nil + } + + // ensure that custom field names are valid + // no leading or trailing whitespace, no empty strings + for k := range values { + if err := s.validateCustomFieldName(k); err != nil { + return fmt.Errorf("custom field name %q: %w", k, err) + } + } + + // ensure delete keys are not also in values + for _, k := range deleteKeys { + if _, ok := values[k]; ok { + return fmt.Errorf("custom field name %q cannot be in both values and delete keys", k) + } + } + + return nil +} + +func (s *customFieldsStore) validateCustomFieldName(fieldName string) error { + // ensure that custom field names are valid + // no leading or trailing whitespace, no empty strings + if strings.TrimSpace(fieldName) == "" { + return fmt.Errorf("custom field name cannot be empty") + } + if fieldName != strings.TrimSpace(fieldName) { + return fmt.Errorf("custom field name cannot have leading or trailing whitespace") + } + if len(fieldName) > maxCustomFieldNameLength { + return fmt.Errorf("custom field name must be less than %d characters", maxCustomFieldNameLength+1) + } + return nil +} + +func getSQLValueFromCustomFieldInput(input any) (interface{}, error) { + jsonBytes, err := json.Marshal(input) + if err != nil { + return nil, fmt.Errorf("failed to marshal custom field value: %w", err) + } + return string(jsonBytes), nil +} + +func getSQLTypeFromInput(input any) string { + return reflect.TypeOf(input).String() +} + +func (s *customFieldsStore) sqlValueToValue(value interface{}, gotype string) (interface{}, error) { + val, ok := value.([]byte) + if !ok { + return value, nil + } + + var res interface{} + + if err := json.Unmarshal(val, &res); err != nil { + return nil, fmt.Errorf("failed to unmarshal JSONB value: %w", err) + } + + switch gotype { + case "string": + str, ok := res.(string) + if !ok { + return nil, fmt.Errorf("expected string, got %T", res) + } + return str, nil + case "int64": + f, ok := res.(float64) + if !ok { + return nil, fmt.Errorf("expected float64 for int64 conversion, got %T", res) + } + return int64(f), nil + case "float64": + f, ok := res.(float64) + if !ok { + return nil, fmt.Errorf("expected float64, got %T", res) + } + return f, nil + case "[]interface {}": + arr, ok := res.([]interface{}) + if !ok { + return nil, fmt.Errorf("expected array, got %T", res) + } + return arr, nil + case "map[string]interface {}": + m, ok := res.(map[string]interface{}) + if !ok { + return nil, fmt.Errorf("expected map, got %T", res) + } + return m, nil + default: + return res, nil + } +} + +func (s *customFieldsStore) setCustomFields(ctx context.Context, id int, values map[string]interface{}, partial bool) error { + if !partial { + // delete existing custom fields + if err := s.deleteForID(ctx, id); err != nil { + return err + } + } + + if len(values) == 0 { + return nil + } + + conflictKey := s.fk.GetCol().(string) + ", field" + // upsert new custom fields + q := dialect.Insert(s.table).Prepared(true).Cols(s.fk, "field", "value", "type"). + OnConflict(goqu.DoUpdate(conflictKey, goqu.Record{"value": goqu.I("excluded.value"), "type": goqu.I("excluded.type")})) + r := make([]interface{}, len(values)) + var i int + for key, value := range values { + v, err := getSQLValueFromCustomFieldInput(value) + if err != nil { + return fmt.Errorf("getting SQL value for field %q: %w", key, err) + } + r[i] = goqu.Record{ + "field": key, + "value": goqu.L("?::JSONB", v), + "type": getSQLTypeFromInput(value), + s.fk.GetCol().(string): id, + } + i++ + } + + if _, err := exec(ctx, q.Rows(r...)); err != nil { + return fmt.Errorf("inserting custom fields: %w", err) + } + + return nil +} + +func (s *customFieldsStore) deleteCustomFields(ctx context.Context, id int, keys []string) error { + if len(keys) == 0 { + return nil + } + + q := dialect.Delete(s.table). + Where(s.fk.Eq(id)). + Where(goqu.I("field").In(keys)) + + if _, err := exec(ctx, q); err != nil { + return fmt.Errorf("deleting custom fields: %w", err) + } + + return nil +} + +func (s *customFieldsStore) GetCustomFields(ctx context.Context, id int) (map[string]interface{}, error) { + q := dialect.Select("field", "value", "type").From(s.table).Where(s.fk.Eq(id)) + + const single = false + ret := make(map[string]interface{}) + err := queryFunc(ctx, q, single, func(rows *sqlx.Rows) error { + var err error + var field string + var value interface{} + var gotype string + if err = rows.Scan(&field, &value, &gotype); err != nil { + return fmt.Errorf("scanning custom fields: %w", err) + } + ret[field], err = s.sqlValueToValue(value, gotype) + return err + }) + if err != nil { + return nil, fmt.Errorf("getting custom fields: %w", err) + } + + return ret, nil +} + +func (s *customFieldsStore) GetCustomFieldsBulk(ctx context.Context, ids []int) ([]models.CustomFieldMap, error) { + q := dialect.Select(s.fk.As("id"), "field", "value", "type").From(s.table).Where(s.fk.In(ids)) + + const single = false + ret := make([]models.CustomFieldMap, len(ids)) + + idi := make(map[int]int, len(ids)) + for i, id := range ids { + idi[id] = i + } + + err := queryFunc(ctx, q, single, func(rows *sqlx.Rows) error { + var err error + var id int + var field string + var value interface{} + var gotype string + if err = rows.Scan(&id, &field, &value, &gotype); err != nil { + return fmt.Errorf("scanning custom fields: %w", err) + } + + i := idi[id] + m := ret[i] + if m == nil { + m = make(map[string]interface{}) + ret[i] = m + } + + m[field], err = s.sqlValueToValue(value, gotype) + return err + }) + if err != nil { + return nil, fmt.Errorf("getting custom fields: %w", err) + } + + return ret, nil +} + +type customFieldsFilterHandler struct { + table string + fkCol string + c []models.CustomFieldCriterionInput + idCol string +} + +func (h *customFieldsFilterHandler) innerJoin(f *filterBuilder, as string, field string) { + joinOn := fmt.Sprintf("%s = %s.%s AND %s.field = ?", h.idCol, as, h.fkCol, as) + f.addInnerJoin(h.table, as, joinOn, field) +} + +func (h *customFieldsFilterHandler) leftJoin(f *filterBuilder, as string, field string) { + joinOn := fmt.Sprintf("%s = %s.%s AND %s.field = ?", h.idCol, as, h.fkCol, as) + f.addLeftJoin(h.table, as, joinOn, field) +} + +func (h *customFieldsFilterHandler) handleCriterion(f *filterBuilder, joinAs string, cc models.CustomFieldCriterionInput) { + // convert values + cv := append([]interface{}{}, cc.Value...) + + valueAsString := fmt.Sprintf("%s.value->>0", joinAs) + valueAsNumber := fmt.Sprintf("%s.value::numeric", joinAs) + valueIsNumber := fmt.Sprintf("jsonb_typeof(%s.value) = 'number'", joinAs) + + switch cc.Modifier { + case models.CriterionModifierEquals: + h.innerJoin(f, joinAs, cc.Field) + f.addWhere(fmt.Sprintf("%s ILIKE %s", valueAsString, getInBinding(len(cv))), cv...) + case models.CriterionModifierNotEquals: + h.innerJoin(f, joinAs, cc.Field) + f.addWhere(fmt.Sprintf("%s NOT ILIKE %s", valueAsString, getInBinding(len(cv))), cv...) + case models.CriterionModifierIncludes: + clauses := make([]sqlClause, len(cv)) + for i, v := range cv { + clauses[i] = makeClause(fmt.Sprintf("%s ILIKE ?", valueAsString), fmt.Sprintf("%%%v%%", v)) + } + h.innerJoin(f, joinAs, cc.Field) + f.whereClauses = append(f.whereClauses, clauses...) + case models.CriterionModifierExcludes: + for _, v := range cv { + f.addWhere(fmt.Sprintf("%s NOT ILIKE ?", valueAsString), fmt.Sprintf("%%%v%%", v)) + } + h.leftJoin(f, joinAs, cc.Field) + case models.CriterionModifierMatchesRegex: + for _, v := range cv { + vs, ok := v.(string) + if !ok { + f.setError(fmt.Errorf("unsupported custom field criterion value type: %T", v)) + } + if _, err := regexp.Compile(vs); err != nil { + f.setError(err) + return + } + f.addWhere(fmt.Sprintf("regex_match(%s, ?)", valueAsString), v) + } + h.innerJoin(f, joinAs, cc.Field) + case models.CriterionModifierNotMatchesRegex: + for _, v := range cv { + vs, ok := v.(string) + if !ok { + f.setError(fmt.Errorf("unsupported custom field criterion value type: %T", v)) + } + if _, err := regexp.Compile(vs); err != nil { + f.setError(err) + return + } + f.addWhere(fmt.Sprintf("(%s.value IS NULL OR NOT regex_match(%s, ?))", joinAs, valueAsString), v) + } + h.leftJoin(f, joinAs, cc.Field) + case models.CriterionModifierIsNull: + h.leftJoin(f, joinAs, cc.Field) + f.addWhere(fmt.Sprintf("%s.value IS NULL OR TRIM(%s) = ''", joinAs, valueAsString)) + case models.CriterionModifierNotNull: + h.innerJoin(f, joinAs, cc.Field) + f.addWhere(fmt.Sprintf("TRIM(%s) != ''", valueAsString)) + case models.CriterionModifierBetween: + if len(cv) != 2 { + f.setError(fmt.Errorf("expected 2 values for custom field criterion modifier BETWEEN, got %d", len(cv))) + return + } + h.innerJoin(f, joinAs, cc.Field) + f.addWhere(valueIsNumber) + f.addWhere(fmt.Sprintf("%s BETWEEN ? AND ?", valueAsNumber), cv[0], cv[1]) + case models.CriterionModifierNotBetween: + h.innerJoin(f, joinAs, cc.Field) + f.addWhere(valueIsNumber) + f.addWhere(fmt.Sprintf("%s NOT BETWEEN ? AND ?", valueAsNumber), cv[0], cv[1]) + case models.CriterionModifierLessThan: + if len(cv) != 1 { + f.setError(fmt.Errorf("expected 1 value for custom field criterion modifier LESS_THAN, got %d", len(cv))) + return + } + h.innerJoin(f, joinAs, cc.Field) + f.addWhere(valueIsNumber) + f.addWhere(fmt.Sprintf("%s < ?", valueAsNumber), cv[0]) + case models.CriterionModifierGreaterThan: + if len(cv) != 1 { + f.setError(fmt.Errorf("expected 1 value for custom field criterion modifier LESS_THAN, got %d", len(cv))) + return + } + h.innerJoin(f, joinAs, cc.Field) + f.addWhere(valueIsNumber) + f.addWhere(fmt.Sprintf("%s > ?", valueAsNumber), cv[0]) + default: + f.setError(fmt.Errorf("unsupported custom field criterion modifier: %s", cc.Modifier)) + } +} + +func (h *customFieldsFilterHandler) handle(ctx context.Context, f *filterBuilder) { + if len(h.c) == 0 { + return + } + + for i, cc := range h.c { + join := fmt.Sprintf("custom_fields_%d", i) + h.handleCriterion(f, join, cc) + } +} diff --git a/pkg/postgres/custom_migrations.go b/pkg/postgres/custom_migrations.go new file mode 100644 index 0000000000..221dad9e5d --- /dev/null +++ b/pkg/postgres/custom_migrations.go @@ -0,0 +1,24 @@ +package postgres + +import ( + "context" + + "github.com/jmoiron/sqlx" +) + +type customMigrationFunc func(ctx context.Context, db *sqlx.DB) error + +func RegisterPostMigration(schemaVersion uint, fn customMigrationFunc) { + v := postMigrations[schemaVersion] + v = append(v, fn) + postMigrations[schemaVersion] = v +} + +func RegisterPreMigration(schemaVersion uint, fn customMigrationFunc) { + v := preMigrations[schemaVersion] + v = append(v, fn) + preMigrations[schemaVersion] = v +} + +var postMigrations = make(map[uint][]customMigrationFunc) +var preMigrations = make(map[uint][]customMigrationFunc) diff --git a/pkg/postgres/database.go b/pkg/postgres/database.go new file mode 100644 index 0000000000..aa015772bd --- /dev/null +++ b/pkg/postgres/database.go @@ -0,0 +1,438 @@ +package postgres + +import ( + "context" + "database/sql" + "embed" + "errors" + "fmt" + "path/filepath" + "runtime" + "strings" + "time" + + _ "github.com/doug-martin/goqu/v9/dialect/postgres" + _ "github.com/jackc/pgx/v5/stdlib" + "github.com/jmoiron/sqlx" + + "github.com/stashapp/stash/pkg/database" + "github.com/stashapp/stash/pkg/logger" +) + +const ( + // TODO: Test for optimality + maxWriteConnections = 5 + maxReadConnections = 15 + // Idle connection timeout, in seconds + // Closes a connection after a period of inactivity, which saves on memory and + // causes the sqlite -wal and -shm files to be automatically deleted. + dbConnTimeout = 30 * time.Second +) + +var appSchemaVersion uint = 13 + +//go:embed migrations/*.sql +var migrationsBox embed.FS + +type storeRepository struct { + Blobs *BlobStore + File *FileStore + Folder *FolderStore + Image *ImageStore + Gallery *GalleryStore + GalleryChapter *GalleryChapterStore + Scene *SceneStore + SceneMarker *SceneMarkerStore + Performer *PerformerStore + SavedFilter *SavedFilterStore + Studio *StudioStore + Tag *TagStore + Group *GroupStore +} + +type Database struct { + *storeRepository + + readDB *sqlx.DB + writeDB *sqlx.DB + dbPath string + + schemaVersion uint + + lockChan chan struct{} +} + +func NewDatabase() *Database { + fileStore := NewFileStore() + folderStore := NewFolderStore() + galleryStore := NewGalleryStore(fileStore, folderStore) + blobStore := NewBlobStore(database.BlobStoreOptions{}) + performerStore := NewPerformerStore(blobStore) + studioStore := NewStudioStore(blobStore) + tagStore := NewTagStore(blobStore) + + r := &storeRepository{} + *r = storeRepository{ + Blobs: blobStore, + File: fileStore, + Folder: folderStore, + Scene: NewSceneStore(r, blobStore), + SceneMarker: NewSceneMarkerStore(), + Image: NewImageStore(r), + Gallery: galleryStore, + GalleryChapter: NewGalleryChapterStore(), + Performer: performerStore, + Studio: studioStore, + Tag: tagStore, + Group: NewGroupStore(blobStore), + SavedFilter: NewSavedFilterStore(), + } + + ret := &Database{ + storeRepository: r, + lockChan: make(chan struct{}, 1), + } + + return ret +} + +func (db *Database) DatabaseBackend() database.DatabaseType { + return database.PostgresBackend +} + +func (db *Database) SetBlobStoreOptions(options database.BlobStoreOptions) { + *db.storeRepository.Blobs = *NewBlobStore(options) +} + +// Ready returns an error if the database is not ready to begin transactions. +func (db *Database) Ready() error { + if db.readDB == nil || db.writeDB == nil { + return database.ErrDatabaseNotInitialized + } + + return nil +} + +// Open initializes the database. If the database is new, then it +// performs a full migration to the latest schema version. Otherwise, any +// necessary migrations must be run separately using RunMigrations. +// Returns true if the database is new. +func (db *Database) Open(dbPath string) error { + db.lock() + defer db.unlock() + + db.dbPath, _ = strings.CutPrefix(dbPath, string(database.PostgresBackend)+":") + + databaseSchemaVersion, err := db.getDatabaseSchemaVersion() + if err != nil { + return fmt.Errorf("getting database schema version: %w", err) + } + + db.schemaVersion = databaseSchemaVersion + + isNew := databaseSchemaVersion == 0 + + if isNew { + // new database, just run the migrations + if err := db.RunAllMigrations(); err != nil { + return fmt.Errorf("error running initial schema migrations: %w", err) + } + } else { + if databaseSchemaVersion > appSchemaVersion { + return &database.MismatchedSchemaVersionError{ + CurrentSchemaVersion: databaseSchemaVersion, + RequiredSchemaVersion: appSchemaVersion, + } + } + + // if migration is needed, then don't open the connection + if db.needsMigration() { + return &database.MigrationNeededError{ + CurrentSchemaVersion: databaseSchemaVersion, + RequiredSchemaVersion: appSchemaVersion, + } + } + } + + if err := db.initialise(); err != nil { + return err + } + + if isNew { + // optimize database after migration + err = db.Optimise(context.Background()) + if err != nil { + logger.Warnf("error while performing post-migration optimisation: %v", err) + } + } + + return nil +} + +// lock locks the database for writing. This method will block until the lock is acquired. +func (db *Database) lock() { + db.lockChan <- struct{}{} +} + +// unlock unlocks the database +func (db *Database) unlock() { + // will block the caller if the lock is not held, so check first + select { + case <-db.lockChan: + return + default: + panic("database is not locked") + } +} + +func (db *Database) Close() error { + db.lock() + defer db.unlock() + + if db.readDB != nil { + if err := db.readDB.Close(); err != nil { + return err + } + + db.readDB = nil + } + if db.writeDB != nil { + if err := db.writeDB.Close(); err != nil { + return err + } + + db.writeDB = nil + } + + return nil +} + +func (db *Database) open(disableForeignKeys bool, writable bool) (conn *sqlx.DB, err error) { + conn, err = sqlx.Open("pgx", db.dbPath) + + if err != nil { + return nil, fmt.Errorf("db.Open(): %w", err) + } + + if disableForeignKeys { + _, err = conn.Exec("SET session_replication_role = replica;") + + if err != nil { + return nil, fmt.Errorf("conn.Exec(): %w", err) + } + } + if !writable { + _, err = conn.Exec("SET SESSION CHARACTERISTICS AS TRANSACTION READ ONLY;") + + if err != nil { + return nil, fmt.Errorf("conn.Exec(): %w", err) + } + } + + return conn, nil +} + +func (db *Database) initialise() error { + if err := db.openReadDB(); err != nil { + return fmt.Errorf("opening read database: %w", err) + } + if err := db.openWriteDB(); err != nil { + return fmt.Errorf("opening write database: %w", err) + } + + return nil +} + +func (db *Database) openReadDB() error { + const ( + disableForeignKeys = false + writable = false + ) + var err error + db.readDB, err = db.open(disableForeignKeys, writable) + db.readDB.SetMaxOpenConns(maxReadConnections) + db.readDB.SetMaxIdleConns(maxReadConnections) + db.readDB.SetConnMaxIdleTime(dbConnTimeout) + return err +} + +func (db *Database) openWriteDB() error { + const ( + disableForeignKeys = false + writable = true + ) + var err error + db.writeDB, err = db.open(disableForeignKeys, writable) + db.writeDB.SetMaxOpenConns(maxWriteConnections) + db.writeDB.SetMaxIdleConns(maxWriteConnections) + db.writeDB.SetConnMaxIdleTime(dbConnTimeout) + return err +} + +func (db *Database) Remove() (err error) { + _, err = db.writeDB.Exec(` +DO $$ +DECLARE + r record; +BEGIN + FOR r IN SELECT quote_ident(tablename) AS tablename, quote_ident(schemaname) AS schemaname FROM pg_tables WHERE schemaname = 'public' + LOOP + RAISE INFO 'Dropping table %.%', r.schemaname, r.tablename; + EXECUTE format('DROP TABLE IF EXISTS %I.%I CASCADE', r.schemaname, r.tablename); + END LOOP; +END$$; +`) + + return err +} + +func (db *Database) Reset() error { + databasePath := db.dbPath + if err := db.Remove(); err != nil { + return err + } + + if err := db.Open(databasePath); err != nil { + return fmt.Errorf("[reset DB] unable to initialize: %w", err) + } + + return nil +} + +// Backup the database. If db is nil, then uses the existing database +// connection. +func (db *Database) Backup(backupPath string) (err error) { + logger.Warn("Postgres backend detected, ignoring Backup request") + return nil +} + +func (db *Database) Anonymise(outPath string) error { + anon, err := NewAnonymiser(db, outPath) + + if err != nil { + return err + } + + return anon.Anonymise(context.Background()) +} + +func (db *Database) RestoreFromBackup(backupPath string) error { + logger.Warn("Postgres backend detected, ignoring RestoreFromBackup request") + return nil +} + +func (db *Database) AppSchemaVersion() uint { + return appSchemaVersion +} + +func (db *Database) DatabasePath() string { + return db.dbPath +} + +func (db *Database) DatabaseBackupPath(backupDirectoryPath string) string { + logger.Warn("Postgres backend detected, ignoring DatabaseBackupPath request") + return "" +} + +func (db *Database) AnonymousDatabasePath(backupDirectoryPath string) string { + fn := fmt.Sprintf("%s.anonymous.%d.%s", "postgres", db.schemaVersion, time.Now().Format("20060102_150405")) + + if backupDirectoryPath != "" { + return filepath.Join(backupDirectoryPath, fn) + } + + return fn +} + +func (db *Database) Version() uint { + return db.schemaVersion +} + +func (db *Database) Optimise(ctx context.Context) error { + logger.Info("Optimising database") + + err := db.Analyze(ctx) + if err != nil { + return fmt.Errorf("performing optimization: %w", err) + } + + err = db.Vacuum(ctx) + if err != nil { + return fmt.Errorf("performing vacuum: %w", err) + } + + return nil +} + +// Vacuum runs a VACUUM on the database, rebuilding the database file into a minimal amount of disk space. +func (db *Database) Vacuum(ctx context.Context) error { + _, err := db.writeDB.ExecContext(ctx, "VACUUM (FULL, ANALYZE, VERBOSE)") + return err +} + +// Analyze runs an ANALYZE on the database to improve query performance. +func (db *Database) Analyze(ctx context.Context) error { + return analyze(ctx, db.writeDB) +} + +// analyze runs an ANALYZE on the database to improve query performance. +func analyze(ctx context.Context, db *sqlx.DB) error { + _, err := db.ExecContext(ctx, "ANALYZE") + return err +} + +func (db *Database) ExecSQL(ctx context.Context, query string, args []interface{}) (*int64, *int64, error) { + wrapper := dbWrapperType{} + + result, err := wrapper.Exec(ctx, query, args...) + if err != nil { + return nil, nil, err + } + + var rowsAffected *int64 + ra, err := result.RowsAffected() + if err == nil { + rowsAffected = &ra + } + + return rowsAffected, nil, nil +} + +func (db *Database) QuerySQL(ctx context.Context, query string, args []interface{}) ([]string, [][]interface{}, error) { + wrapper := dbWrapperType{} + + rows, err := wrapper.QueryxContext(ctx, query, args...) + if err != nil && !errors.Is(err, sql.ErrNoRows) { + return nil, nil, err + } + defer rows.Close() + + cols, err := rows.Columns() + if err != nil { + return nil, nil, err + } + + var ret [][]interface{} + + for rows.Next() { + row, err := rows.SliceScan() + if err != nil { + return nil, nil, err + } + ret = append(ret, row) + } + + if err := rows.Err(); err != nil { + return nil, nil, err + } + + return cols, ret, nil +} + +func getBasenameSQL(sql string) string { + if runtime.GOOS == "windows" { + return fmt.Sprintf("basename(%s, '\\')", sql) + } + + return fmt.Sprintf("basename(%s)", sql) +} diff --git a/pkg/postgres/date.go b/pkg/postgres/date.go new file mode 100644 index 0000000000..2dc573da79 --- /dev/null +++ b/pkg/postgres/date.go @@ -0,0 +1,19 @@ +package postgres + +import ( + "github.com/stashapp/stash/pkg/database" + "github.com/stashapp/stash/pkg/models" + "gopkg.in/guregu/null.v4" +) + +// Date represents a date stored as "YYYY-MM-DD" +type Date = database.Date + +// NullDate represents a nullable date stored as "YYYY-MM-DD" +type NullDate = database.NullDate + +var NullDateFromDatePtr = database.NullDateFromDatePtr + +func datePrecisionFromDatePtr(d *models.Date) null.Int { + return database.DatePrecisionFromDatePtr(d) +} diff --git a/pkg/postgres/doc.go b/pkg/postgres/doc.go new file mode 100644 index 0000000000..812b382db4 --- /dev/null +++ b/pkg/postgres/doc.go @@ -0,0 +1,2 @@ +// Package postgres provides interfaces to interact with the postgres database. +package postgres diff --git a/pkg/postgres/file.go b/pkg/postgres/file.go new file mode 100644 index 0000000000..bc63ef24a3 --- /dev/null +++ b/pkg/postgres/file.go @@ -0,0 +1,1042 @@ +package postgres + +import ( + "context" + "database/sql" + "errors" + "fmt" + "io/fs" + "path/filepath" + "strings" + "time" + + "github.com/doug-martin/goqu/v9" + "github.com/doug-martin/goqu/v9/exp" + "github.com/jmoiron/sqlx" + "github.com/stashapp/stash/pkg/models" + "gopkg.in/guregu/null.v4" +) + +const ( + fileTable = "files" + videoFileTable = "video_files" + imageFileTable = "image_files" + fileIDColumn = "file_id" + + videoCaptionsTable = "video_captions" + captionCodeColumn = "language_code" + captionFilenameColumn = "filename" + captionTypeColumn = "caption_type" +) + +type basicFileRow struct { + ID models.FileID `db:"id" goqu:"skipinsert"` + Basename string `db:"basename"` + ZipFileID null.Int `db:"zip_file_id"` + ParentFolderID models.FolderID `db:"parent_folder_id"` + Size int64 `db:"size"` + ModTime Timestamp `db:"mod_time"` + CreatedAt Timestamp `db:"created_at"` + UpdatedAt Timestamp `db:"updated_at"` +} + +func (r *basicFileRow) fromBasicFile(o models.BaseFile) { + r.ID = o.ID + r.Basename = o.Basename + r.ZipFileID = nullIntFromFileIDPtr(o.ZipFileID) + r.ParentFolderID = o.ParentFolderID + r.Size = o.Size + r.ModTime = Timestamp{Timestamp: o.ModTime} + r.CreatedAt = Timestamp{Timestamp: o.CreatedAt} + r.UpdatedAt = Timestamp{Timestamp: o.UpdatedAt} +} + +type videoFileRow struct { + FileID models.FileID `db:"file_id"` + Format string `db:"format"` + Width int `db:"width"` + Height int `db:"height"` + Duration float64 `db:"duration"` + VideoCodec string `db:"video_codec"` + AudioCodec string `db:"audio_codec"` + FrameRate float64 `db:"frame_rate"` + BitRate int64 `db:"bit_rate"` + Interactive bool `db:"interactive"` + InteractiveSpeed null.Int `db:"interactive_speed"` +} + +func (f *videoFileRow) fromVideoFile(ff models.VideoFile) { + f.FileID = ff.ID + f.Format = ff.Format + f.Width = ff.Width + f.Height = ff.Height + f.Duration = ff.Duration + f.VideoCodec = ff.VideoCodec + f.AudioCodec = ff.AudioCodec + f.FrameRate = ff.FrameRate + f.BitRate = ff.BitRate + f.Interactive = ff.Interactive + f.InteractiveSpeed = intFromPtr(ff.InteractiveSpeed) +} + +type imageFileRow struct { + FileID models.FileID `db:"file_id"` + Format string `db:"format"` + Width int `db:"width"` + Height int `db:"height"` +} + +func (f *imageFileRow) fromImageFile(ff models.ImageFile) { + f.FileID = ff.ID + f.Format = ff.Format + f.Width = ff.Width + f.Height = ff.Height +} + +// we redefine this to change the columns around +// otherwise, we collide with the image file columns +type videoFileQueryRow struct { + FileID null.Int `db:"file_id_video"` + Format null.String `db:"video_format"` + Width null.Int `db:"video_width"` + Height null.Int `db:"video_height"` + Duration null.Float `db:"duration"` + VideoCodec null.String `db:"video_codec"` + AudioCodec null.String `db:"audio_codec"` + FrameRate null.Float `db:"frame_rate"` + BitRate null.Int `db:"bit_rate"` + Interactive null.Bool `db:"interactive"` + InteractiveSpeed null.Int `db:"interactive_speed"` +} + +func (f *videoFileQueryRow) resolve() *models.VideoFile { + return &models.VideoFile{ + Format: f.Format.String, + Width: int(f.Width.Int64), + Height: int(f.Height.Int64), + Duration: f.Duration.Float64, + VideoCodec: f.VideoCodec.String, + AudioCodec: f.AudioCodec.String, + FrameRate: f.FrameRate.Float64, + BitRate: f.BitRate.Int64, + Interactive: f.Interactive.Bool, + InteractiveSpeed: nullIntPtr(f.InteractiveSpeed), + } +} + +func videoFileQueryColumns() []interface{} { + table := videoFileTableMgr.table + return []interface{}{ + table.Col("file_id").As("file_id_video"), + table.Col("format").As("video_format"), + table.Col("width").As("video_width"), + table.Col("height").As("video_height"), + table.Col("duration"), + table.Col("video_codec"), + table.Col("audio_codec"), + table.Col("frame_rate"), + table.Col("bit_rate"), + table.Col("interactive"), + table.Col("interactive_speed"), + } +} + +// we redefine this to change the columns around +// otherwise, we collide with the video file columns +type imageFileQueryRow struct { + Format null.String `db:"image_format"` + Width null.Int `db:"image_width"` + Height null.Int `db:"image_height"` +} + +func (imageFileQueryRow) columns(table *table) []interface{} { + ex := table.table + return []interface{}{ + ex.Col("format").As("image_format"), + ex.Col("width").As("image_width"), + ex.Col("height").As("image_height"), + } +} + +func (f *imageFileQueryRow) resolve() *models.ImageFile { + return &models.ImageFile{ + Format: f.Format.String, + Width: int(f.Width.Int64), + Height: int(f.Height.Int64), + } +} + +type fileQueryRow struct { + FileID null.Int `db:"file_id"` + Basename null.String `db:"basename"` + ZipFileID null.Int `db:"zip_file_id"` + ParentFolderID null.Int `db:"parent_folder_id"` + Size null.Int `db:"size"` + ModTime NullTimestamp `db:"mod_time"` + CreatedAt NullTimestamp `db:"file_created_at"` + UpdatedAt NullTimestamp `db:"file_updated_at"` + + ZipBasename null.String `db:"zip_basename"` + ZipFolderPath null.String `db:"zip_folder_path"` + ZipSize null.Int `db:"zip_size"` + + FolderPath null.String `db:"parent_folder_path"` + fingerprintQueryRow + videoFileQueryRow + imageFileQueryRow +} + +func (r *fileQueryRow) resolve() models.File { + basic := &models.BaseFile{ + ID: models.FileID(r.FileID.Int64), + DirEntry: models.DirEntry{ + ZipFileID: nullIntFileIDPtr(r.ZipFileID), + ModTime: r.ModTime.Timestamp.UTC(), + }, + Path: filepath.Join(r.FolderPath.String, r.Basename.String), + ParentFolderID: models.FolderID(r.ParentFolderID.Int64), + Basename: r.Basename.String, + Size: r.Size.Int64, + CreatedAt: r.CreatedAt.Timestamp.UTC(), + UpdatedAt: r.UpdatedAt.Timestamp.UTC(), + } + + if basic.ZipFileID != nil && r.ZipFolderPath.Valid && r.ZipBasename.Valid { + basic.ZipFile = &models.BaseFile{ + ID: *basic.ZipFileID, + Path: filepath.Join(r.ZipFolderPath.String, r.ZipBasename.String), + Basename: r.ZipBasename.String, + Size: r.ZipSize.Int64, + } + } + + var ret models.File = basic + + if r.videoFileQueryRow.Format.Valid { + vf := r.videoFileQueryRow.resolve() + vf.BaseFile = basic + ret = vf + } + + if r.imageFileQueryRow.Format.Valid { + imf := r.imageFileQueryRow.resolve() + imf.BaseFile = basic + ret = imf + } + + r.appendRelationships(basic) + + return ret +} + +func appendFingerprintsUnique(vs []models.Fingerprint, v ...models.Fingerprint) []models.Fingerprint { + for _, vv := range v { + found := false + for _, vsv := range vs { + if vsv.Type == vv.Type { + found = true + break + } + } + + if !found { + vs = append(vs, vv) + } + } + return vs +} + +func (r *fileQueryRow) appendRelationships(i *models.BaseFile) { + if r.fingerprintQueryRow.valid() { + i.Fingerprints = appendFingerprintsUnique(i.Fingerprints, r.fingerprintQueryRow.resolve()) + } +} + +type fileQueryRows []fileQueryRow + +func (r fileQueryRows) resolve() []models.File { + var ret []models.File + var last models.File + var lastID models.FileID + + for _, row := range r { + if last == nil || lastID != models.FileID(row.FileID.Int64) { + f := row.resolve() + last = f + lastID = models.FileID(row.FileID.Int64) + ret = append(ret, last) + continue + } + + // must be merging with previous row + row.appendRelationships(last.Base()) + } + + return ret +} + +type fileRepositoryType struct { + repository + scenes joinRepository + images joinRepository + galleries joinRepository +} + +var ( + fileRepository = fileRepositoryType{ + repository: repository{ + tableName: fileTable, + idColumn: idColumn, + }, + scenes: joinRepository{ + repository: repository{ + tableName: scenesFilesTable, + idColumn: fileIDColumn, + }, + fkColumn: sceneIDColumn, + }, + images: joinRepository{ + repository: repository{ + tableName: imagesFilesTable, + idColumn: fileIDColumn, + }, + fkColumn: imageIDColumn, + }, + galleries: joinRepository{ + repository: repository{ + tableName: galleriesFilesTable, + idColumn: fileIDColumn, + }, + fkColumn: galleryIDColumn, + }, + } +) + +type FileStore struct { + repository + + tableMgr *table +} + +func NewFileStore() *FileStore { + return &FileStore{ + repository: repository{ + tableName: fileTable, + idColumn: idColumn, + }, + + tableMgr: fileTableMgr, + } +} + +func (qb *FileStore) table() exp.IdentifierExpression { + return qb.tableMgr.table +} + +func (qb *FileStore) Create(ctx context.Context, f models.File) error { + var r basicFileRow + r.fromBasicFile(*f.Base()) + + id, err := qb.tableMgr.insertID(ctx, r) + if err != nil { + return err + } + + fileID := models.FileID(id) + + // create extended stuff here + switch ef := f.(type) { + case *models.VideoFile: + if err := qb.createVideoFile(ctx, fileID, *ef); err != nil { + return err + } + case *models.ImageFile: + if err := qb.createImageFile(ctx, fileID, *ef); err != nil { + return err + } + } + + if err := FingerprintReaderWriter.insertJoins(ctx, fileID, f.Base().Fingerprints); err != nil { + return err + } + + updated, err := qb.Find(ctx, fileID) + if err != nil { + return fmt.Errorf("finding after create: %w", err) + } + + base := f.Base() + *base = *updated[0].Base() + + return nil +} + +func (qb *FileStore) Update(ctx context.Context, f models.File) error { + var r basicFileRow + r.fromBasicFile(*f.Base()) + + id := f.Base().ID + + if err := qb.tableMgr.updateByID(ctx, id, r); err != nil { + return err + } + + // create extended stuff here + switch ef := f.(type) { + case *models.VideoFile: + if err := qb.updateOrCreateVideoFile(ctx, id, *ef); err != nil { + return err + } + case *models.ImageFile: + if err := qb.updateOrCreateImageFile(ctx, id, *ef); err != nil { + return err + } + } + + if err := FingerprintReaderWriter.replaceJoins(ctx, id, f.Base().Fingerprints); err != nil { + return err + } + + return nil +} + +// ModifyFingerprints updates existing fingerprints and adds new ones. +func (qb *FileStore) ModifyFingerprints(ctx context.Context, fileID models.FileID, fingerprints []models.Fingerprint) error { + return FingerprintReaderWriter.upsertJoins(ctx, fileID, fingerprints) +} + +func (qb *FileStore) DestroyFingerprints(ctx context.Context, fileID models.FileID, types []string) error { + return FingerprintReaderWriter.destroyJoins(ctx, fileID, types) +} + +func (qb *FileStore) Destroy(ctx context.Context, id models.FileID) error { + return qb.tableMgr.destroyExisting(ctx, []int{int(id)}) +} + +func (qb *FileStore) createVideoFile(ctx context.Context, id models.FileID, f models.VideoFile) error { + var r videoFileRow + r.fromVideoFile(f) + r.FileID = id + if _, err := videoFileTableMgr.insert(ctx, r); err != nil { + return err + } + + return nil +} + +func (qb *FileStore) updateOrCreateVideoFile(ctx context.Context, id models.FileID, f models.VideoFile) error { + exists, err := videoFileTableMgr.idExists(ctx, id) + if err != nil { + return err + } + + if !exists { + return qb.createVideoFile(ctx, id, f) + } + + var r videoFileRow + r.fromVideoFile(f) + r.FileID = id + if err := videoFileTableMgr.updateByID(ctx, id, r); err != nil { + return err + } + + return nil +} + +func (qb *FileStore) createImageFile(ctx context.Context, id models.FileID, f models.ImageFile) error { + var r imageFileRow + r.fromImageFile(f) + r.FileID = id + if _, err := imageFileTableMgr.insert(ctx, r); err != nil { + return err + } + + return nil +} + +func (qb *FileStore) updateOrCreateImageFile(ctx context.Context, id models.FileID, f models.ImageFile) error { + exists, err := imageFileTableMgr.idExists(ctx, id) + if err != nil { + return err + } + + if !exists { + return qb.createImageFile(ctx, id, f) + } + + var r imageFileRow + r.fromImageFile(f) + r.FileID = id + if err := imageFileTableMgr.updateByID(ctx, id, r); err != nil { + return err + } + + return nil +} + +func (qb *FileStore) selectDataset() *goqu.SelectDataset { + table := qb.table() + + folderTable := folderTableMgr.table + fingerprintTable := fingerprintTableMgr.table + videoFileTable := videoFileTableMgr.table + imageFileTable := imageFileTableMgr.table + + zipFileTable := table.As("zip_files") + zipFolderTable := folderTable.As("zip_files_folders") + + cols := []interface{}{ + table.Col("id").As("file_id"), + table.Col("basename"), + table.Col("zip_file_id"), + table.Col("parent_folder_id"), + table.Col("size"), + table.Col("mod_time"), + table.Col("created_at").As("file_created_at"), + table.Col("updated_at").As("file_updated_at"), + folderTable.Col("path").As("parent_folder_path"), + fingerprintTable.Col("type").As("fingerprint_type"), + fingerprintTable.Col("fingerprint"), + zipFileTable.Col("basename").As("zip_basename"), + zipFolderTable.Col("path").As("zip_folder_path"), + // size is needed to open containing zip files + zipFileTable.Col("size").As("zip_size"), + } + + cols = append(cols, videoFileQueryColumns()...) + cols = append(cols, imageFileQueryRow{}.columns(imageFileTableMgr)...) + + ret := dialect.From(table).Select(cols...) + + return ret.InnerJoin( + folderTable, + goqu.On(table.Col("parent_folder_id").Eq(folderTable.Col(idColumn))), + ).LeftJoin( + fingerprintTable, + goqu.On(table.Col(idColumn).Eq(fingerprintTable.Col(fileIDColumn))), + ).LeftJoin( + videoFileTable, + goqu.On(table.Col(idColumn).Eq(videoFileTable.Col(fileIDColumn))), + ).LeftJoin( + imageFileTable, + goqu.On(table.Col(idColumn).Eq(imageFileTable.Col(fileIDColumn))), + ).LeftJoin( + zipFileTable, + goqu.On(table.Col("zip_file_id").Eq(zipFileTable.Col("id"))), + ).LeftJoin( + zipFolderTable, + goqu.On(zipFileTable.Col("parent_folder_id").Eq(zipFolderTable.Col(idColumn))), + ) +} + +func (qb *FileStore) countDataset() *goqu.SelectDataset { + table := qb.table() + + folderTable := folderTableMgr.table + fingerprintTable := fingerprintTableMgr.table + videoFileTable := videoFileTableMgr.table + imageFileTable := imageFileTableMgr.table + + zipFileTable := table.As("zip_files") + zipFolderTable := folderTable.As("zip_files_folders") + + ret := dialect.From(table).Select(goqu.COUNT(goqu.DISTINCT(table.Col("id")))) + + return ret.InnerJoin( + folderTable, + goqu.On(table.Col("parent_folder_id").Eq(folderTable.Col(idColumn))), + ).LeftJoin( + fingerprintTable, + goqu.On(table.Col(idColumn).Eq(fingerprintTable.Col(fileIDColumn))), + ).LeftJoin( + videoFileTable, + goqu.On(table.Col(idColumn).Eq(videoFileTable.Col(fileIDColumn))), + ).LeftJoin( + imageFileTable, + goqu.On(table.Col(idColumn).Eq(imageFileTable.Col(fileIDColumn))), + ).LeftJoin( + zipFileTable, + goqu.On(table.Col("zip_file_id").Eq(zipFileTable.Col("id"))), + ).LeftJoin( + zipFolderTable, + goqu.On(zipFileTable.Col("parent_folder_id").Eq(zipFolderTable.Col(idColumn))), + ) +} + +func (qb *FileStore) get(ctx context.Context, q *goqu.SelectDataset) (models.File, error) { + ret, err := qb.getMany(ctx, q) + if err != nil { + return nil, err + } + + if len(ret) == 0 { + return nil, sql.ErrNoRows + } + + return ret[0], nil +} + +func (qb *FileStore) getMany(ctx context.Context, q *goqu.SelectDataset) ([]models.File, error) { + const single = false + var rows fileQueryRows + if err := queryFunc(ctx, q, single, func(r *sqlx.Rows) error { + var f fileQueryRow + if err := r.StructScan(&f); err != nil { + return err + } + + f.fingerprintQueryRow.correct() + + rows = append(rows, f) + return nil + }); err != nil { + return nil, err + } + + return rows.resolve(), nil +} + +func (qb *FileStore) Find(ctx context.Context, ids ...models.FileID) ([]models.File, error) { + var files []models.File + for _, id := range ids { + file, err := qb.find(ctx, id) + if err != nil { + return nil, err + } + + if file == nil { + return nil, fmt.Errorf("file with id %d not found", id) + } + + files = append(files, file) + } + + return files, nil +} + +func (qb *FileStore) find(ctx context.Context, id models.FileID) (models.File, error) { + q := qb.selectDataset().Where(qb.tableMgr.byID(id)) + + ret, err := qb.get(ctx, q) + if err != nil { + return nil, fmt.Errorf("getting file by id %d: %w", id, err) + } + + return ret, nil +} + +// FindByPath returns the first file that matches the given path. Wildcard characters are supported. +func (qb *FileStore) FindByPath(ctx context.Context, p string, caseSensitive bool) (models.File, error) { + + ret, err := qb.FindAllByPath(ctx, p, caseSensitive) + + if err != nil { + return nil, err + } + + if len(ret) == 0 { + return nil, nil + } + + return ret[0], nil +} + +// FindAllByPath returns all the files that match the given path. +// Wildcard characters are supported. +func (qb *FileStore) FindAllByPath(ctx context.Context, p string, caseSensitive bool) ([]models.File, error) { + // separate basename from path + basename := filepath.Base(p) + dirName := filepath.Dir(p) + + // replace wildcards + basename = strings.ReplaceAll(basename, "*", "%") + dirName = strings.ReplaceAll(dirName, "*", "%") + + table := qb.table() + folderTable := folderTableMgr.table + + // like uses case-insensitive matching. Only use like if wildcards are used + q := qb.selectDataset().Prepared(true) + + if strings.Contains(basename, "%") || strings.Contains(dirName, "%") || !caseSensitive { + q = q.Where( + folderTable.Col("path").ILike(dirName), + table.Col("basename").ILike(basename), + ) + } else { + q = q.Where( + folderTable.Col("path").Eq(dirName), + table.Col("basename").Eq(basename), + ) + } + + ret, err := qb.getMany(ctx, q) + if err != nil && !errors.Is(err, sql.ErrNoRows) { + return nil, fmt.Errorf("getting file by path %s: %w", p, err) + } + + return ret, nil +} + +func (qb *FileStore) allInPaths(q *goqu.SelectDataset, p []string) *goqu.SelectDataset { + folderTable := folderTableMgr.table + + var conds []exp.Expression + for _, pp := range p { + ppWildcard := pp + string(filepath.Separator) + "%" + + conds = append(conds, folderTable.Col("path").Eq(pp), folderTable.Col("path").ILike(ppWildcard)) + } + + return q.Where( + goqu.Or(conds...), + ) +} + +// FindAllByPaths returns the all files that are within any of the given paths. +// Returns all if limit is < 0. +// Returns all files if p is empty. +func (qb *FileStore) FindAllInPaths(ctx context.Context, p []string, limit, offset int) ([]models.File, error) { + table := qb.table() + folderTable := folderTableMgr.table + + q := dialect.From(table).Prepared(true).InnerJoin( + folderTable, + goqu.On(table.Col("parent_folder_id").Eq(folderTable.Col(idColumn))), + ).Select(table.Col(idColumn)) + + q = qb.allInPaths(q, p) + + if limit > -1 { + q = q.Limit(uint(limit)) + } + + q = q.Offset(uint(offset)) + + ret, err := qb.findBySubquery(ctx, q) + if err != nil && !errors.Is(err, sql.ErrNoRows) { + return nil, fmt.Errorf("getting files by path %s: %w", p, err) + } + + return ret, nil +} + +// CountAllInPaths returns a count of all files that are within any of the given paths. +// Returns count of all files if p is empty. +func (qb *FileStore) CountAllInPaths(ctx context.Context, p []string) (int, error) { + q := qb.countDataset().Prepared(true) + q = qb.allInPaths(q, p) + + return count(ctx, q) +} + +func (qb *FileStore) findBySubquery(ctx context.Context, sq *goqu.SelectDataset) ([]models.File, error) { + table := qb.table() + + q := qb.selectDataset().Prepared(true).Where( + table.Col(idColumn).Eq( + sq, + ), + ) + + return qb.getMany(ctx, q) +} + +func (qb *FileStore) FindByFingerprint(ctx context.Context, fp models.Fingerprint) ([]models.File, error) { + fingerprintTable := fingerprintTableMgr.table + + fingerprints := fingerprintTable.As("fp") + + sq := dialect.From(fingerprints).Select(fingerprints.Col(fileIDColumn)).Where( + fingerprints.Col("type").Eq(fp.Type), + fingerprints.Col("fingerprint").Eq(fp.Fingerprint), + ) + + return qb.findBySubquery(ctx, sq) +} + +func (qb *FileStore) FindByZipFileID(ctx context.Context, zipFileID models.FileID) ([]models.File, error) { + table := qb.table() + + q := qb.selectDataset().Prepared(true).Where( + table.Col("zip_file_id").Eq(zipFileID), + ) + + return qb.getMany(ctx, q) +} + +// FindByFileInfo finds files that match the base name, size, and mod time of the given file. +func (qb *FileStore) FindByFileInfo(ctx context.Context, info fs.FileInfo, size int64) ([]models.File, error) { + table := qb.table() + + modTime := info.ModTime().Format(time.RFC3339) + + q := qb.selectDataset().Prepared(true).Where( + table.Col("basename").Eq(info.Name()), + table.Col("size").Eq(size), + table.Col("mod_time").Eq(modTime), + ) + + return qb.getMany(ctx, q) +} + +func (qb *FileStore) CountByFolderID(ctx context.Context, folderID models.FolderID) (int, error) { + table := qb.table() + + q := qb.countDataset().Prepared(true).Where( + table.Col("parent_folder_id").Eq(folderID), + ) + + return count(ctx, q) +} + +func (qb *FileStore) IsPrimary(ctx context.Context, fileID models.FileID) (bool, error) { + joinTables := []exp.IdentifierExpression{ + scenesFilesJoinTable, + galleriesFilesJoinTable, + imagesFilesJoinTable, + } + + var sq *goqu.SelectDataset + + for _, t := range joinTables { + qq := dialect.From(t).Select(t.Col(fileIDColumn)).Where( + t.Col(fileIDColumn).Eq(fileID), + t.Col("primary").IsTrue(), + ) + + if sq == nil { + sq = qq + } else { + sq = sq.Union(qq) + } + } + + q := dialect.Select(goqu.COUNT("*").As("count")).Prepared(true).From( + sq, + ) + + var ret int + if err := querySimple(ctx, q, &ret); err != nil { + return false, err + } + + return ret > 0, nil +} + +func (qb *FileStore) validateFilter(fileFilter *models.FileFilterType) error { + const and = "AND" + const or = "OR" + const not = "NOT" + + if fileFilter.And != nil { + if fileFilter.Or != nil { + return illegalFilterCombination(and, or) + } + if fileFilter.Not != nil { + return illegalFilterCombination(and, not) + } + + return qb.validateFilter(fileFilter.And) + } + + if fileFilter.Or != nil { + if fileFilter.Not != nil { + return illegalFilterCombination(or, not) + } + + return qb.validateFilter(fileFilter.Or) + } + + if fileFilter.Not != nil { + return qb.validateFilter(fileFilter.Not) + } + + return nil +} + +func (qb *FileStore) makeFilter(ctx context.Context, fileFilter *models.FileFilterType) *filterBuilder { + query := &filterBuilder{} + + if fileFilter.And != nil { + query.and(qb.makeFilter(ctx, fileFilter.And)) + } + if fileFilter.Or != nil { + query.or(qb.makeFilter(ctx, fileFilter.Or)) + } + if fileFilter.Not != nil { + query.not(qb.makeFilter(ctx, fileFilter.Not)) + } + + filter := filterBuilderFromHandler(ctx, &fileFilterHandler{ + fileFilter: fileFilter, + }) + + return filter +} + +func (qb *FileStore) Query(ctx context.Context, options models.FileQueryOptions) (*models.FileQueryResult, error) { + fileFilter := options.FileFilter + findFilter := options.FindFilter + + if fileFilter == nil { + fileFilter = &models.FileFilterType{} + } + if findFilter == nil { + findFilter = &models.FindFilterType{} + } + + query := qb.newQuery() + query.join(folderTable, "", "files.parent_folder_id = folders.id") + + distinctIDs(&query, fileTable) + + if q := findFilter.Q; q != nil && *q != "" { + filepathColumn := "folders.path || '" + string(filepath.Separator) + "' || files.basename" + searchColumns := []string{filepathColumn} + query.parseQueryString(searchColumns, *q) + } + + if err := qb.validateFilter(fileFilter); err != nil { + return nil, err + } + filter := qb.makeFilter(ctx, fileFilter) + + if err := query.addFilter(filter); err != nil { + return nil, err + } + + if err := qb.setQuerySort(&query, findFilter); err != nil { + return nil, err + } + query.pagination = getPagination(findFilter) + + result, err := qb.queryGroupedFields(ctx, options, query) + if err != nil { + return nil, fmt.Errorf("error querying aggregate fields: %w", err) + } + + idsResult, err := query.findIDs(ctx) + if err != nil { + return nil, fmt.Errorf("error finding IDs: %w", err) + } + + result.IDs = make([]models.FileID, len(idsResult)) + for i, id := range idsResult { + result.IDs[i] = models.FileID(id) + } + + return result, nil +} + +func (qb *FileStore) queryGroupedFields(ctx context.Context, options models.FileQueryOptions, query queryBuilder) (*models.FileQueryResult, error) { + if !options.Count && !options.TotalDuration && !options.Megapixels && !options.TotalSize { + // nothing to do - return empty result + return models.NewFileQueryResult(qb), nil + } + + aggregateQuery := qb.newQuery() + + if options.Count { + aggregateQuery.addColumn("COUNT(DISTINCT temp.id) as total") + } + + if options.TotalDuration { + query.addJoins( + join{ + table: videoFileTable, + onClause: "files.id = video_files.file_id", + }, + ) + query.addColumn("COALESCE(video_files.duration, 0) as duration") + aggregateQuery.addColumn("COALESCE(SUM(temp.duration), 0) as duration") + } + if options.Megapixels { + query.addJoins( + join{ + table: imageFileTable, + onClause: "files.id = image_files.file_id", + }, + ) + query.addColumn("COALESCE(image_files.width, 0) * COALESCE(image_files.height, 0) as megapixels") + aggregateQuery.addColumn("COALESCE(SUM(temp.megapixels), 0) / 1000000 as megapixels") + } + + if options.TotalSize { + query.addColumn("COALESCE(files.size, 0) as size") + aggregateQuery.addColumn("COALESCE(SUM(temp.size), 0) as size") + } + + const includeSortPagination = false + aggregateQuery.from = fmt.Sprintf("(%s) as temp", query.toSQL(includeSortPagination)) + + out := struct { + Total int + Duration float64 + Megapixels float64 + Size int64 + }{} + if err := qb.repository.queryStruct(ctx, aggregateQuery.toSQL(includeSortPagination), query.args, &out); err != nil { + return nil, err + } + + ret := models.NewFileQueryResult(qb) + ret.Count = out.Total + ret.Megapixels = out.Megapixels + ret.TotalDuration = out.Duration + ret.TotalSize = out.Size + + return ret, nil +} + +var fileSortOptions = sortOptions{ + "created_at", + "id", + "path", + "random", + "updated_at", +} + +func (qb *FileStore) setQuerySort(query *queryBuilder, findFilter *models.FindFilterType) error { + if findFilter == nil || findFilter.Sort == nil || *findFilter.Sort == "" { + return nil + } + sort := findFilter.GetSort("path") + + // CVE-2024-32231 - ensure sort is in the list of allowed sorts + if err := fileSortOptions.validateSort(sort); err != nil { + return err + } + + direction := findFilter.GetDirection() + switch sort { + case "path": + // special handling for path + query.sort += fmt.Sprintf(" ORDER BY folders.path %s, files.basename %[1]s", direction) + query.addGroupBy("folders.path", "files.basename") + default: + add, agg := getSort(sort, direction, "files") + query.sort += add + query.addGroupBy(agg...) + } + + return nil +} + +func (qb *FileStore) captionRepository() *captionRepository { + return &captionRepository{ + repository: repository{ + tableName: videoCaptionsTable, + idColumn: fileIDColumn, + }, + } +} + +func (qb *FileStore) GetCaptions(ctx context.Context, fileID models.FileID) ([]*models.VideoCaption, error) { + return qb.captionRepository().get(ctx, fileID) +} + +func (qb *FileStore) UpdateCaptions(ctx context.Context, fileID models.FileID, captions []*models.VideoCaption) error { + return qb.captionRepository().replace(ctx, fileID, captions) +} diff --git a/pkg/postgres/file_filter.go b/pkg/postgres/file_filter.go new file mode 100644 index 0000000000..461638e62b --- /dev/null +++ b/pkg/postgres/file_filter.go @@ -0,0 +1,346 @@ +package postgres + +import ( + "context" + "fmt" + + "github.com/stashapp/stash/pkg/models" + "github.com/stashapp/stash/pkg/utils" +) + +type fileFilterHandler struct { + fileFilter *models.FileFilterType + // if true, don't allow use of related filters + isRelated bool +} + +func (qb *fileFilterHandler) validate() error { + fileFilter := qb.fileFilter + if fileFilter == nil { + return nil + } + + if err := validateFilterCombination(fileFilter.OperatorFilter); err != nil { + return err + } + + if qb.isRelated && (fileFilter.ScenesFilter != nil || fileFilter.ImagesFilter != nil || fileFilter.GalleriesFilter != nil) { + return fmt.Errorf("cannot use related filters inside a related filter") + } + + if subFilter := fileFilter.SubFilter(); subFilter != nil { + sqb := &fileFilterHandler{fileFilter: subFilter, isRelated: qb.isRelated} + if err := sqb.validate(); err != nil { + return err + } + } + + return nil +} + +func (qb *fileFilterHandler) handle(ctx context.Context, f *filterBuilder) { + fileFilter := qb.fileFilter + if fileFilter == nil { + return + } + + if err := qb.validate(); err != nil { + f.setError(err) + return + } + + sf := fileFilter.SubFilter() + if sf != nil { + sub := &fileFilterHandler{sf, qb.isRelated} + handleSubFilter(ctx, sub, f, fileFilter.OperatorFilter) + } + + f.handleCriterion(ctx, qb.criterionHandler()) +} + +func (qb *fileFilterHandler) criterionHandler() criterionHandler { + fileFilter := qb.fileFilter + return compoundHandler{ + &videoFileFilterHandler{ + filter: fileFilter.VideoFileFilter, + }, + &imageFileFilterHandler{ + filter: fileFilter.ImageFileFilter, + }, + + pathCriterionHandler(fileFilter.Path, "folders.path", "files.basename", nil), + stringCriterionHandler(fileFilter.Basename, "files.basename"), + stringCriterionHandler(fileFilter.Dir, "folders.path"), + ×tampCriterionHandler{fileFilter.ModTime, "files.mod_time", nil}, + + qb.parentFolderCriterionHandler(fileFilter.ParentFolder), + qb.zipFileCriterionHandler(fileFilter.ZipFile), + + qb.sceneCountCriterionHandler(fileFilter.SceneCount), + qb.imageCountCriterionHandler(fileFilter.ImageCount), + qb.galleryCountCriterionHandler(fileFilter.GalleryCount), + + qb.hashesCriterionHandler(fileFilter.Hashes), + + qb.phashDuplicatedCriterionHandler(fileFilter.Duplicated), + ×tampCriterionHandler{fileFilter.CreatedAt, "files.created_at", nil}, + ×tampCriterionHandler{fileFilter.UpdatedAt, "files.updated_at", nil}, + + &relatedFilterHandler{ + relatedIDCol: "scenes_files.scene_id", + relatedRepo: sceneRepository.repository, + relatedHandler: &sceneFilterHandler{fileFilter.ScenesFilter}, + joinFn: func(f *filterBuilder) { + fileRepository.scenes.innerJoin(f, "", "files.id") + }, + }, + &relatedFilterHandler{ + relatedIDCol: "images_files.image_id", + relatedRepo: imageRepository.repository, + relatedHandler: &imageFilterHandler{fileFilter.ImagesFilter}, + joinFn: func(f *filterBuilder) { + fileRepository.images.innerJoin(f, "", "files.id") + }, + }, + &relatedFilterHandler{ + relatedIDCol: "galleries_files.gallery_id", + relatedRepo: galleryRepository.repository, + relatedHandler: &galleryFilterHandler{fileFilter.GalleriesFilter}, + joinFn: func(f *filterBuilder) { + fileRepository.galleries.innerJoin(f, "", "files.id") + }, + }, + } +} + +func (qb *fileFilterHandler) zipFileCriterionHandler(criterion *models.MultiCriterionInput) criterionHandlerFunc { + return func(ctx context.Context, f *filterBuilder) { + if criterion != nil { + if criterion.Modifier == models.CriterionModifierIsNull || criterion.Modifier == models.CriterionModifierNotNull { + var notClause string + if criterion.Modifier == models.CriterionModifierNotNull { + notClause = "NOT" + } + + f.addWhere(fmt.Sprintf("files.zip_file_id IS %s NULL", notClause)) + return + } + + if len(criterion.Value) == 0 { + return + } + + var args []interface{} + for _, tagID := range criterion.Value { + args = append(args, tagID) + } + + whereClause := "" + havingClause := "" + switch criterion.Modifier { + case models.CriterionModifierIncludes: + whereClause = "files.zip_file_id IN " + getInBinding(len(criterion.Value)) + case models.CriterionModifierExcludes: + whereClause = "files.zip_file_id NOT IN " + getInBinding(len(criterion.Value)) + } + + f.addWhere(whereClause, args...) + f.addHaving(havingClause) + } + } +} + +func (qb *fileFilterHandler) parentFolderCriterionHandler(folder *models.HierarchicalMultiCriterionInput) criterionHandlerFunc { + return func(ctx context.Context, f *filterBuilder) { + if folder == nil { + return + } + + folderCopy := *folder + switch folderCopy.Modifier { + case models.CriterionModifierEquals: + folderCopy.Modifier = models.CriterionModifierIncludesAll + case models.CriterionModifierNotEquals: + folderCopy.Modifier = models.CriterionModifierExcludes + } + + hh := hierarchicalMultiCriterionHandlerBuilder{ + primaryTable: fileTable, + foreignTable: folderTable, + foreignFK: "parent_folder_id", + parentFK: "parent_folder_id", + } + + hh.handler(&folderCopy)(ctx, f) + } +} + +func (qb *fileFilterHandler) sceneCountCriterionHandler(c *models.IntCriterionInput) criterionHandlerFunc { + h := countCriterionHandlerBuilder{ + primaryTable: fileTable, + joinTable: scenesFilesTable, + primaryFK: fileIDColumn, + } + + return h.handler(c) +} + +func (qb *fileFilterHandler) imageCountCriterionHandler(c *models.IntCriterionInput) criterionHandlerFunc { + h := countCriterionHandlerBuilder{ + primaryTable: fileTable, + joinTable: imagesFilesTable, + primaryFK: fileIDColumn, + } + + return h.handler(c) +} + +func (qb *fileFilterHandler) galleryCountCriterionHandler(c *models.IntCriterionInput) criterionHandlerFunc { + h := countCriterionHandlerBuilder{ + primaryTable: fileTable, + joinTable: galleriesFilesTable, + primaryFK: fileIDColumn, + } + + return h.handler(c) +} + +func (qb *fileFilterHandler) phashDuplicatedCriterionHandler(duplicatedFilter *models.PHashDuplicationCriterionInput) criterionHandlerFunc { + return func(ctx context.Context, f *filterBuilder) { + // TODO: Wishlist item: Implement Distance matching + if duplicatedFilter != nil { + var v string + if *duplicatedFilter.Duplicated { + v = ">" + } else { + v = "=" + } + + f.addInnerJoin("(SELECT file_id FROM files_fingerprints INNER JOIN (SELECT fingerprint FROM files_fingerprints WHERE type = 'phash' GROUP BY fingerprint HAVING COUNT (fingerprint) "+v+" 1) dupes on files_fingerprints.fingerprint = dupes.fingerprint)", "scph", "files.id = scph.file_id") + } + } +} + +func (qb *fileFilterHandler) hashesCriterionHandler(hashes []*models.FingerprintFilterInput) criterionHandlerFunc { + return func(ctx context.Context, f *filterBuilder) { + // TODO - this won't work for AND/OR combinations + for i, hash := range hashes { + t := fmt.Sprintf("file_fingerprints_%d", i) + f.addLeftJoin(fingerprintTable, t, fmt.Sprintf("files.id = %s.file_id AND %s.type = ?", t, t), hash.Type) + + value, _ := utils.StringToPhash(hash.Value) + distance := 0 + if hash.Distance != nil { + distance = *hash.Distance + } + + if distance > 0 { + // needed to avoid a type mismatch + f.addWhere(fmt.Sprintf("typeof(%s.fingerprint) = 'integer'", t)) + f.addWhere(fmt.Sprintf("phash_distance(%s.fingerprint, ?) < ?", t), value, distance) + } else { + // use the default handler + intCriterionHandler(&models.IntCriterionInput{ + Value: int(value), + Modifier: models.CriterionModifierEquals, + }, t+".fingerprint", nil)(ctx, f) + } + } + } +} + +type videoFileFilterHandler struct { + filter *models.VideoFileFilterInput +} + +func (qb *videoFileFilterHandler) handle(ctx context.Context, f *filterBuilder) { + videoFileFilter := qb.filter + if videoFileFilter == nil { + return + } + f.handleCriterion(ctx, qb.criterionHandler()) +} + +func (qb *videoFileFilterHandler) criterionHandler() criterionHandler { + videoFileFilter := qb.filter + return compoundHandler{ + joinedStringCriterionHandler(videoFileFilter.Format, "video_files.format", qb.addVideoFilesTable), + floatIntCriterionHandler(videoFileFilter.Duration, "video_files.duration", qb.addVideoFilesTable), + resolutionCriterionHandler(videoFileFilter.Resolution, "video_files.height", "video_files.width", qb.addVideoFilesTable), + orientationCriterionHandler(videoFileFilter.Orientation, "video_files.height", "video_files.width", qb.addVideoFilesTable), + floatIntCriterionHandler(videoFileFilter.Framerate, "ROUND(video_files.frame_rate)", qb.addVideoFilesTable), + intCriterionHandler(videoFileFilter.Bitrate, "video_files.bit_rate", qb.addVideoFilesTable), + qb.codecCriterionHandler(videoFileFilter.VideoCodec, "video_files.video_codec", qb.addVideoFilesTable), + qb.codecCriterionHandler(videoFileFilter.AudioCodec, "video_files.audio_codec", qb.addVideoFilesTable), + + boolCriterionHandler(videoFileFilter.Interactive, "video_files.interactive", qb.addVideoFilesTable), + intCriterionHandler(videoFileFilter.InteractiveSpeed, "video_files.interactive_speed", qb.addVideoFilesTable), + + qb.captionCriterionHandler(videoFileFilter.Captions), + } +} + +func (qb *videoFileFilterHandler) addVideoFilesTable(f *filterBuilder) { + f.addLeftJoin(videoFileTable, "", "video_files.file_id = files.id") +} + +func (qb *videoFileFilterHandler) codecCriterionHandler(codec *models.StringCriterionInput, codecColumn string, addJoinFn func(f *filterBuilder)) criterionHandlerFunc { + return func(ctx context.Context, f *filterBuilder) { + if codec != nil { + if addJoinFn != nil { + addJoinFn(f) + } + + stringCriterionHandler(codec, codecColumn)(ctx, f) + } + } +} + +func (qb *videoFileFilterHandler) captionCriterionHandler(captions *models.StringCriterionInput) criterionHandlerFunc { + h := stringListCriterionHandlerBuilder{ + primaryTable: sceneTable, + primaryFK: sceneIDColumn, + joinTable: videoCaptionsTable, + stringColumn: captionCodeColumn, + addJoinTable: func(f *filterBuilder) { + f.addLeftJoin(videoCaptionsTable, "", "video_captions.file_id = files.id") + }, + excludeHandler: func(f *filterBuilder, criterion *models.StringCriterionInput) { + excludeClause := `files.id NOT IN ( + SELECT files.id from files + INNER JOIN video_captions on video_captions.file_id = files.id + WHERE video_captions.language_code LIKE ? + )` + f.addWhere(excludeClause, criterion.Value) + + // TODO - should we also exclude null values? + }, + } + + return h.handler(captions) +} + +type imageFileFilterHandler struct { + filter *models.ImageFileFilterInput +} + +func (qb *imageFileFilterHandler) handle(ctx context.Context, f *filterBuilder) { + ff := qb.filter + if ff == nil { + return + } + f.handleCriterion(ctx, qb.criterionHandler()) +} + +func (qb *imageFileFilterHandler) criterionHandler() criterionHandler { + ff := qb.filter + return compoundHandler{ + joinedStringCriterionHandler(ff.Format, "image_files.format", qb.addImageFilesTable), + resolutionCriterionHandler(ff.Resolution, "image_files.height", "image_files.width", qb.addImageFilesTable), + orientationCriterionHandler(ff.Orientation, "image_files.height", "image_files.width", qb.addImageFilesTable), + } +} + +func (qb *imageFileFilterHandler) addImageFilesTable(f *filterBuilder) { + f.addLeftJoin(imageFileTable, "", "image_files.file_id = files.id") +} diff --git a/pkg/postgres/filter.go b/pkg/postgres/filter.go new file mode 100644 index 0000000000..61b8ec339b --- /dev/null +++ b/pkg/postgres/filter.go @@ -0,0 +1,447 @@ +package postgres + +import ( + "context" + "errors" + "fmt" + "strings" + + "github.com/stashapp/stash/pkg/models" +) + +func illegalFilterCombination(type1, type2 string) error { + return fmt.Errorf("cannot have %s and %s in the same filter", type1, type2) +} + +func validateFilterCombination[T any](sf models.OperatorFilter[T]) error { + const and = "AND" + const or = "OR" + const not = "NOT" + + if sf.And != nil { + if sf.Or != nil { + return illegalFilterCombination(and, or) + } + if sf.Not != nil { + return illegalFilterCombination(and, not) + } + } + + if sf.Or != nil { + if sf.Not != nil { + return illegalFilterCombination(or, not) + } + } + + return nil +} + +func handleSubFilter[T any](ctx context.Context, handler criterionHandler, f *filterBuilder, subFilter models.OperatorFilter[T]) { + subQuery := &filterBuilder{} + handler.handle(ctx, subQuery) + + if subFilter.And != nil { + f.and(subQuery) + } + if subFilter.Or != nil { + f.or(subQuery) + } + if subFilter.Not != nil { + f.not(subQuery) + } +} + +type sqlClause struct { + sql string + args []interface{} +} + +func (c sqlClause) not() sqlClause { + return sqlClause{ + sql: "NOT (" + c.sql + ")", + args: c.args, + } +} + +func makeClause(sql string, args ...interface{}) sqlClause { + return sqlClause{ + sql: sql, + args: args, + } +} + +func joinClauses(joinType string, clauses ...sqlClause) sqlClause { + var ret []string + var args []interface{} + + for _, clause := range clauses { + ret = append(ret, "("+clause.sql+")") + args = append(args, clause.args...) + } + + return sqlClause{sql: strings.Join(ret, " "+joinType+" "), args: args} +} + +func orClauses(clauses ...sqlClause) sqlClause { + return joinClauses("OR", clauses...) +} + +func andClauses(clauses ...sqlClause) sqlClause { + return joinClauses("AND", clauses...) +} + +type join struct { + table string + as string + onClause string + joinType string + args []interface{} + + // if true, indicates this is required for sorting only + sort bool +} + +// equals returns true if the other join alias/table is equal to this one +func (j join) equals(o join) bool { + return j.alias() == o.alias() +} + +// alias returns the as string, or the table if as is empty +func (j join) alias() string { + if j.as == "" { + return j.table + } + + return j.as +} + +func (j join) toSQL() string { + asStr := "" + joinStr := j.joinType + if j.as != "" && j.as != j.table { + asStr = " AS " + j.as + } + if j.joinType == "" { + joinStr = "LEFT" + } + + return fmt.Sprintf("%s JOIN %s%s ON %s", joinStr, j.table, asStr, j.onClause) +} + +type joins []join + +// addUnique only adds if not already present +// returns true if added +func (j *joins) addUnique(newJoin join) bool { + found := false + for i, jj := range *j { + if jj.equals(newJoin) { + found = true + // if sort is false on the new join, but true on the existing, set the false + if !newJoin.sort && jj.sort { + (*j)[i].sort = false + } + break + } + } + + if !found { + *j = append(*j, newJoin) + } + return !found +} + +func (j *joins) add(newJoins ...join) { + // only add if not already joined + for _, newJoin := range newJoins { + j.addUnique(newJoin) + } +} + +func (j *joins) toSQL(includeSortPagination bool) string { + if len(*j) == 0 { + return "" + } + + var ret []string + for _, jj := range *j { + // skip sort-only joins if not including sort/pagination + if !includeSortPagination && jj.sort { + continue + } + ret = append(ret, jj.toSQL()) + } + + return " " + strings.Join(ret, " ") +} + +type filterBuilder struct { + subFilter *filterBuilder + subFilterOp string + + joins joins + whereClauses []sqlClause + havingClauses []sqlClause + withClauses []sqlClause + recursiveWith bool + + err error +} + +func (f *filterBuilder) empty() bool { + return f == nil || (len(f.whereClauses) == 0 && len(f.joins) == 0 && len(f.havingClauses) == 0 && f.subFilter == nil) +} + +func filterBuilderFromHandler(ctx context.Context, handler criterionHandler) *filterBuilder { + f := &filterBuilder{} + handler.handle(ctx, f) + return f +} + +var errSubFilterAlreadySet = errors.New(`sub-filter already set`) + +// sub-filter operator values +var ( + andOp = "AND" + orOp = "OR" + notOp = "AND NOT" +) + +// and sets the sub-filter that will be ANDed with this one. +// Sets the error state if sub-filter is already set. +func (f *filterBuilder) and(a *filterBuilder) { + if f.subFilter != nil { + f.setError(errSubFilterAlreadySet) + return + } + + f.subFilter = a + f.subFilterOp = andOp +} + +// or sets the sub-filter that will be ORed with this one. +// Sets the error state if a sub-filter is already set. +func (f *filterBuilder) or(o *filterBuilder) { + if f.subFilter != nil { + f.setError(errSubFilterAlreadySet) + return + } + + f.subFilter = o + f.subFilterOp = orOp +} + +// not sets the sub-filter that will be AND NOTed with this one. +// Sets the error state if a sub-filter is already set. +func (f *filterBuilder) not(n *filterBuilder) { + if f.subFilter != nil { + f.setError(errSubFilterAlreadySet) + return + } + + f.subFilter = n + f.subFilterOp = notOp +} + +// addLeftJoin adds a left join to the filter. The join is expressed in SQL as: +// LEFT JOIN [AS ] ON +// The AS is omitted if as is empty. +// This method does not add a join if it its alias/table name is already +// present in another existing join. +func (f *filterBuilder) addLeftJoin(table, as, onClause string, args ...interface{}) { + newJoin := join{ + table: table, + as: as, + onClause: onClause, + joinType: "LEFT", + args: args, + } + + f.joins.add(newJoin) +} + +// addInnerJoin adds an inner join to the filter. The join is expressed in SQL as: +// INNER JOIN
[AS ] ON +// The AS is omitted if as is empty. +// This method does not add a join if it its alias/table name is already +// present in another existing join. +func (f *filterBuilder) addInnerJoin(table, as, onClause string, args ...interface{}) { + newJoin := join{ + table: table, + as: as, + onClause: onClause, + joinType: "INNER", + args: args, + } + + f.joins.add(newJoin) +} + +// addWhere adds a where clause and arguments to the filter. Where clauses +// are ANDed together. Does not add anything if the provided string is empty. +func (f *filterBuilder) addWhere(sql string, args ...interface{}) { + if sql == "" { + return + } + f.whereClauses = append(f.whereClauses, makeClause(sql, args...)) +} + +// addHaving adds a where clause and arguments to the filter. Having clauses +// are ANDed together. Does not add anything if the provided string is empty. +func (f *filterBuilder) addHaving(sql string, args ...interface{}) { + if sql == "" { + return + } + f.havingClauses = append(f.havingClauses, makeClause(sql, args...)) +} + +// addWith adds a with clause and arguments to the filter +func (f *filterBuilder) addWith(sql string, args ...interface{}) { + if sql == "" { + return + } + + f.withClauses = append(f.withClauses, makeClause(sql, args...)) +} + +// addRecursiveWith adds a with clause and arguments to the filter, and sets it to recursive +// +//nolint:unused +func (f *filterBuilder) addRecursiveWith(sql string, args ...interface{}) { + if sql == "" { + return + } + + f.addWith(sql, args...) + f.recursiveWith = true +} + +func (f *filterBuilder) getSubFilterClause(clause, subFilterClause string) string { + ret := clause + + if subFilterClause != "" { + var op string + if len(ret) > 0 { + op = " " + f.subFilterOp + " " + } else if f.subFilterOp == notOp { + op = "NOT " + } + + ret += op + "(" + subFilterClause + ")" + } + + return ret +} + +// generateWhereClauses generates the SQL where clause for this filter. +// All where clauses within the filter are ANDed together. This is combined +// with the sub-filter, which will use the applicable operator (AND/OR/AND NOT). +func (f *filterBuilder) generateWhereClauses() (clause string, args []interface{}) { + clause, args = f.andClauses(f.whereClauses) + + if f.subFilter != nil { + c, a := f.subFilter.generateWhereClauses() + if c != "" { + clause = f.getSubFilterClause(clause, c) + if len(a) > 0 { + args = append(args, a...) + } + } + } + + return +} + +// generateHavingClauses generates the SQL having clause for this filter. +// All having clauses within the filter are ANDed together. This is combined +// with the sub-filter, which will use the applicable operator (AND/OR/AND NOT). +func (f *filterBuilder) generateHavingClauses() (string, []interface{}) { + clause, args := f.andClauses(f.havingClauses) + + if f.subFilter != nil { + c, a := f.subFilter.generateHavingClauses() + if c != "" { + clause = f.getSubFilterClause(clause, c) + if len(a) > 0 { + args = append(args, a...) + } + } + } + + return clause, args +} + +func (f *filterBuilder) generateWithClauses() (string, []interface{}) { + var clauses []string + var args []interface{} + for _, w := range f.withClauses { + clauses = append(clauses, w.sql) + args = append(args, w.args...) + } + + if len(clauses) > 0 { + return strings.Join(clauses, ", "), args + } + + return "", nil +} + +// getAllJoins returns all of the joins in this filter and any sub-filter(s). +// Redundant joins will not be duplicated in the return value. +func (f *filterBuilder) getAllJoins() joins { + var ret joins + ret.add(f.joins...) + if f.subFilter != nil { + subJoins := f.subFilter.getAllJoins() + if len(subJoins) > 0 { + ret.add(subJoins...) + } + } + + return ret +} + +// getError returns the error state on this filter, or on any sub-filter(s) if +// the error state is nil. +func (f *filterBuilder) getError() error { + if f.err != nil { + return f.err + } + + if f.subFilter != nil { + return f.subFilter.getError() + } + + return nil +} + +// handleCriterion calls the handle function on the provided criterionHandler, +// providing itself. +func (f *filterBuilder) handleCriterion(ctx context.Context, handler criterionHandler) { + handler.handle(ctx, f) +} + +func (f *filterBuilder) setError(e error) { + if f.err == nil { + f.err = e + } +} + +func (f *filterBuilder) andClauses(input []sqlClause) (string, []interface{}) { + var clauses []string + var args []interface{} + for _, w := range input { + clauses = append(clauses, w.sql) + args = append(args, w.args...) + } + + if len(clauses) > 0 { + c := "(" + strings.Join(clauses, ") AND (") + ")" + if len(clauses) > 1 { + c = "(" + c + ")" + } + return c, args + } + + return "", nil +} diff --git a/pkg/postgres/filter_hierarchical.go b/pkg/postgres/filter_hierarchical.go new file mode 100644 index 0000000000..13d35bb631 --- /dev/null +++ b/pkg/postgres/filter_hierarchical.go @@ -0,0 +1,222 @@ +package postgres + +import ( + "context" + "fmt" + + "github.com/stashapp/stash/pkg/models" + "github.com/stashapp/stash/pkg/utils" +) + +// hierarchicalRelationshipHandler provides handlers for parent, children, parent count, and child count criteria. +type hierarchicalRelationshipHandler struct { + primaryTable string + relationTable string + aliasPrefix string + parentIDCol string + childIDCol string +} + +func (h hierarchicalRelationshipHandler) validateModifier(m models.CriterionModifier) error { + switch m { + case models.CriterionModifierIncludesAll, models.CriterionModifierIncludes, models.CriterionModifierExcludes, models.CriterionModifierIsNull, models.CriterionModifierNotNull: + // valid + return nil + default: + return fmt.Errorf("invalid modifier %s", m) + } +} + +func (h hierarchicalRelationshipHandler) handleNullNotNull(f *filterBuilder, m models.CriterionModifier, isParents bool) { + var notClause string + if m == models.CriterionModifierNotNull { + notClause = "NOT" + } + + as := h.aliasPrefix + "_parents" + col := h.childIDCol + if !isParents { + as = h.aliasPrefix + "_children" + col = h.parentIDCol + } + + // Based on: + // f.addLeftJoin("tags_relations", "parent_relations", "tags.id = parent_relations.child_id") + // f.addWhere(fmt.Sprintf("parent_relations.parent_id IS %s NULL", notClause)) + + f.addLeftJoin(h.relationTable, as, fmt.Sprintf("%s.id = %s.%s", h.primaryTable, as, col)) + f.addWhere(fmt.Sprintf("%s.%s IS %s NULL", as, col, notClause)) +} + +func (h hierarchicalRelationshipHandler) parentsAlias() string { + return h.aliasPrefix + "_parents" +} + +func (h hierarchicalRelationshipHandler) childrenAlias() string { + return h.aliasPrefix + "_children" +} + +func (h hierarchicalRelationshipHandler) valueQuery(value []string, depth int, alias string, isParents bool) string { + var depthCondition string + if depth != -1 { + depthCondition = fmt.Sprintf("WHERE depth < %d", depth) + } + + queryTempl := `{alias} AS ( +SELECT {root_id_col} AS root_id, {item_id_col} AS item_id, 0 AS depth FROM {relation_table} WHERE {root_id_col} IN` + getInBinding(len(value)) + ` +UNION +SELECT root_id, {item_id_col}, depth + 1 FROM {relation_table} INNER JOIN {alias} ON item_id = {root_id_col} ` + depthCondition + ` +)` + + var queryMap utils.StrFormatMap + if isParents { + queryMap = utils.StrFormatMap{ + "root_id_col": h.parentIDCol, + "item_id_col": h.childIDCol, + } + } else { + queryMap = utils.StrFormatMap{ + "root_id_col": h.childIDCol, + "item_id_col": h.parentIDCol, + } + } + + queryMap["alias"] = alias + queryMap["relation_table"] = h.relationTable + + return utils.StrFormat(queryTempl, queryMap) +} + +func (h hierarchicalRelationshipHandler) handleValues(f *filterBuilder, c models.HierarchicalMultiCriterionInput, isParents bool, aliasSuffix string) { + if len(c.Value) == 0 { + return + } + + var args []interface{} + for _, val := range c.Value { + args = append(args, val) + } + + depthVal := 0 + if c.Depth != nil { + depthVal = *c.Depth + } + + tableAlias := h.parentsAlias() + if !isParents { + tableAlias = h.childrenAlias() + } + tableAlias += aliasSuffix + + query := h.valueQuery(c.Value, depthVal, tableAlias, isParents) + f.addRecursiveWith(query, args...) + + f.addLeftJoin(tableAlias, "", fmt.Sprintf("%s.item_id = %s.id", tableAlias, h.primaryTable)) + addHierarchicalConditionClauses(f, c, tableAlias, "root_id") +} + +func (h hierarchicalRelationshipHandler) handleValuesSimple(f *filterBuilder, value string, isParents bool) { + joinCol := h.childIDCol + valueCol := h.parentIDCol + if !isParents { + joinCol = h.parentIDCol + valueCol = h.childIDCol + } + + tableAlias := h.parentsAlias() + if !isParents { + tableAlias = h.childrenAlias() + } + + f.addInnerJoin(h.relationTable, tableAlias, fmt.Sprintf("%s.%s = %s.id", tableAlias, joinCol, h.primaryTable)) + f.addWhere(fmt.Sprintf("%s.%s = ?", tableAlias, valueCol), value) +} + +func (h hierarchicalRelationshipHandler) hierarchicalCriterionHandler(criterion *models.HierarchicalMultiCriterionInput, isParents bool) criterionHandlerFunc { + return func(ctx context.Context, f *filterBuilder) { + if criterion != nil { + c := criterion.CombineExcludes() + + // validate the modifier + if err := h.validateModifier(c.Modifier); err != nil { + f.setError(err) + return + } + + if c.Modifier == models.CriterionModifierIsNull || c.Modifier == models.CriterionModifierNotNull { + h.handleNullNotNull(f, c.Modifier, isParents) + return + } + + if len(c.Value) == 0 && len(c.Excludes) == 0 { + return + } + + depth := 0 + if c.Depth != nil { + depth = *c.Depth + } + + // if we have a single include, no excludes, and no depth, we can use a simple join and where clause + if (c.Modifier == models.CriterionModifierIncludes || c.Modifier == models.CriterionModifierIncludesAll) && len(c.Value) == 1 && len(c.Excludes) == 0 && depth == 0 { + h.handleValuesSimple(f, c.Value[0], isParents) + return + } + + aliasSuffix := "" + h.handleValues(f, c, isParents, aliasSuffix) + + if len(c.Excludes) > 0 { + exCriterion := models.HierarchicalMultiCriterionInput{ + Value: c.Excludes, + Depth: c.Depth, + Modifier: models.CriterionModifierExcludes, + } + + aliasSuffix := "2" + h.handleValues(f, exCriterion, isParents, aliasSuffix) + } + } + } +} + +func (h hierarchicalRelationshipHandler) ParentsCriterionHandler(criterion *models.HierarchicalMultiCriterionInput) criterionHandlerFunc { + const isParents = true + return h.hierarchicalCriterionHandler(criterion, isParents) +} + +func (h hierarchicalRelationshipHandler) ChildrenCriterionHandler(criterion *models.HierarchicalMultiCriterionInput) criterionHandlerFunc { + const isParents = false + return h.hierarchicalCriterionHandler(criterion, isParents) +} + +func (h hierarchicalRelationshipHandler) countCriterionHandler(c *models.IntCriterionInput, isParents bool) criterionHandlerFunc { + tableAlias := h.parentsAlias() + col := h.childIDCol + otherCol := h.parentIDCol + if !isParents { + tableAlias = h.childrenAlias() + col = h.parentIDCol + otherCol = h.childIDCol + } + tableAlias += "_count" + + return func(ctx context.Context, f *filterBuilder) { + if c != nil { + f.addLeftJoin(h.relationTable, tableAlias, fmt.Sprintf("%s.%s = %s.id", tableAlias, col, h.primaryTable)) + clause, args := getIntCriterionWhereClause(fmt.Sprintf("count(distinct %s.%s)", tableAlias, otherCol), *c) + + f.addHaving(clause, args...) + } + } +} + +func (h hierarchicalRelationshipHandler) ParentCountCriterionHandler(parentCount *models.IntCriterionInput) criterionHandlerFunc { + const isParents = true + return h.countCriterionHandler(parentCount, isParents) +} + +func (h hierarchicalRelationshipHandler) ChildCountCriterionHandler(childCount *models.IntCriterionInput) criterionHandlerFunc { + const isParents = false + return h.countCriterionHandler(childCount, isParents) +} diff --git a/pkg/postgres/filter_internal_test.go b/pkg/postgres/filter_internal_test.go new file mode 100644 index 0000000000..460f2c2824 --- /dev/null +++ b/pkg/postgres/filter_internal_test.go @@ -0,0 +1,645 @@ +//go:build pg_integration +// +build pg_integration + +package postgres + +import ( + "context" + "errors" + "fmt" + "testing" + + "github.com/stashapp/stash/pkg/models" + "github.com/stretchr/testify/assert" +) + +var testCtx = context.Background() + +func TestJoinsAddJoin(t *testing.T) { + var joins joins + + // add a single join + joins.add(join{table: "test"}) + + assert := assert.New(t) + + // ensure join was added + assert.Len(joins, 1) + + // add the same join and another + joins.add([]join{ + { + table: "test", + }, + { + table: "foo", + }, + }...) + + // should have added a single join + assert.Len(joins, 2) +} + +func TestFilterBuilderAnd(t *testing.T) { + assert := assert.New(t) + + f := &filterBuilder{} + other := &filterBuilder{} + newBuilder := &filterBuilder{} + + // and should set the subFilter + f.and(other) + assert.Equal(other, f.subFilter) + assert.Nil(f.getError()) + + // and should set error if and is set + f.and(newBuilder) + assert.Equal(other, f.subFilter) + assert.Equal(errSubFilterAlreadySet, f.getError()) + + // and should set error if or is set + // and should not set subFilter if or is set + f = &filterBuilder{} + f.or(other) + f.and(newBuilder) + assert.Equal(other, f.subFilter) + assert.Equal(errSubFilterAlreadySet, f.getError()) + + // and should set error if not is set + // and should not set subFilter if not is set + f = &filterBuilder{} + f.not(other) + f.and(newBuilder) + assert.Equal(other, f.subFilter) + assert.Equal(errSubFilterAlreadySet, f.getError()) +} + +func TestFilterBuilderOr(t *testing.T) { + assert := assert.New(t) + + f := &filterBuilder{} + other := &filterBuilder{} + newBuilder := &filterBuilder{} + + // or should set the orFilter + f.or(other) + assert.Equal(other, f.subFilter) + assert.Nil(f.getError()) + + // or should set error if or is set + f.or(newBuilder) + assert.Equal(newBuilder, f.subFilter) + assert.Equal(errSubFilterAlreadySet, f.getError()) + + // or should set error if and is set + // or should not set subFilter if and is set + f = &filterBuilder{} + f.and(other) + f.or(newBuilder) + assert.Equal(other, f.subFilter) + assert.Equal(errSubFilterAlreadySet, f.getError()) + + // or should set error if not is set + // or should not set subFilter if not is set + f = &filterBuilder{} + f.not(other) + f.or(newBuilder) + assert.Equal(other, f.subFilter) + assert.Equal(errSubFilterAlreadySet, f.getError()) +} + +func TestFilterBuilderNot(t *testing.T) { + assert := assert.New(t) + + f := &filterBuilder{} + other := &filterBuilder{} + newBuilder := &filterBuilder{} + + // not should set the subFilter + f.not(other) + // ensure and filter is set + assert.Equal(other, f.subFilter) + assert.Nil(f.getError()) + + // not should set error if not is set + f.not(newBuilder) + assert.Equal(newBuilder, f.subFilter) + assert.Equal(errSubFilterAlreadySet, f.getError()) + + // not should set error if and is set + // not should not set subFilter if and is set + f = &filterBuilder{} + f.and(other) + f.not(newBuilder) + assert.Equal(other, f.subFilter) + assert.Equal(errSubFilterAlreadySet, f.getError()) + + // not should set error if or is set + // not should not set subFilter if or is set + f = &filterBuilder{} + f.or(other) + f.not(newBuilder) + assert.Equal(other, f.subFilter) + assert.Equal(errSubFilterAlreadySet, f.getError()) +} + +func TestAddJoin(t *testing.T) { + assert := assert.New(t) + + f := &filterBuilder{} + + const ( + table1Name = "table1Name" + table2Name = "table2Name" + + as1Name = "as1" + as2Name = "as2" + + onClause = "onClause1" + ) + + f.addLeftJoin(table1Name, as1Name, onClause) + + // ensure join is added + assert.Len(f.joins, 1) + assert.Equal(fmt.Sprintf("LEFT JOIN %s AS %s ON %s", table1Name, as1Name, onClause), f.joins[0].toSQL()) + + // ensure join with same as is not added + f.addLeftJoin(table2Name, as1Name, onClause) + assert.Len(f.joins, 1) + + // ensure same table with different alias can be added + f.addLeftJoin(table1Name, as2Name, onClause) + assert.Len(f.joins, 2) + assert.Equal(fmt.Sprintf("LEFT JOIN %s AS %s ON %s", table1Name, as2Name, onClause), f.joins[1].toSQL()) + + // ensure table without alias can be added if tableName != existing alias/tableName + f.addLeftJoin(table1Name, "", onClause) + assert.Len(f.joins, 3) + assert.Equal(fmt.Sprintf("LEFT JOIN %s ON %s", table1Name, onClause), f.joins[2].toSQL()) + + // ensure table with alias == table name of a join without alias is not added + f.addLeftJoin(table2Name, table1Name, onClause) + assert.Len(f.joins, 3) + + // ensure table without alias cannot be added if tableName == existing alias + f.addLeftJoin(as2Name, "", onClause) + assert.Len(f.joins, 3) + + // ensure AS is not used if same as table name + f.addLeftJoin(table2Name, table2Name, onClause) + assert.Len(f.joins, 4) + assert.Equal(fmt.Sprintf("LEFT JOIN %s ON %s", table2Name, onClause), f.joins[3].toSQL()) +} + +func TestAddWhere(t *testing.T) { + assert := assert.New(t) + + f := &filterBuilder{} + + // ensure empty sql adds nothing + f.addWhere("") + assert.Len(f.whereClauses, 0) + + const whereClause = "a = b" + var args = []interface{}{"1", "2"} + + // ensure addWhere sets where clause and args + f.addWhere(whereClause, args...) + assert.Len(f.whereClauses, 1) + assert.Equal(whereClause, f.whereClauses[0].sql) + assert.Equal(args, f.whereClauses[0].args) + + // ensure addWhere without args sets where clause + f.addWhere(whereClause) + assert.Len(f.whereClauses, 2) + assert.Equal(whereClause, f.whereClauses[1].sql) + assert.Len(f.whereClauses[1].args, 0) +} + +func TestAddHaving(t *testing.T) { + assert := assert.New(t) + + f := &filterBuilder{} + + // ensure empty sql adds nothing + f.addHaving("") + assert.Len(f.havingClauses, 0) + + const havingClause = "a = b" + var args = []interface{}{"1", "2"} + + // ensure addWhere sets where clause and args + f.addHaving(havingClause, args...) + assert.Len(f.havingClauses, 1) + assert.Equal(havingClause, f.havingClauses[0].sql) + assert.Equal(args, f.havingClauses[0].args) + + // ensure addWhere without args sets where clause + f.addHaving(havingClause) + assert.Len(f.havingClauses, 2) + assert.Equal(havingClause, f.havingClauses[1].sql) + assert.Len(f.havingClauses[1].args, 0) +} + +func TestGenerateWhereClauses(t *testing.T) { + assert := assert.New(t) + + f := &filterBuilder{} + + const clause1 = "a = 1" + const clause2 = "b = 2" + const clause3 = "c = 3" + + const arg1 = "1" + const arg2 = "2" + const arg3 = "3" + + // ensure single where clause is generated correctly + f.addWhere(clause1) + r, rArgs := f.generateWhereClauses() + assert.Equal(fmt.Sprintf("(%s)", clause1), r) + assert.Len(rArgs, 0) + + // ensure multiple where clauses are surrounded with parenthesis and + // ANDed together + f.addWhere(clause2, arg1, arg2) + r, rArgs = f.generateWhereClauses() + assert.Equal(fmt.Sprintf("((%s) AND (%s))", clause1, clause2), r) + assert.Len(rArgs, 2) + + // ensure empty subfilter is not added to generated where clause + sf := &filterBuilder{} + f.and(sf) + + r, rArgs = f.generateWhereClauses() + assert.Equal(fmt.Sprintf("((%s) AND (%s))", clause1, clause2), r) + assert.Len(rArgs, 2) + + // ensure sub-filter is generated correctly + sf.addWhere(clause3, arg3) + r, rArgs = f.generateWhereClauses() + assert.Equal(fmt.Sprintf("((%s) AND (%s)) AND ((%s))", clause1, clause2, clause3), r) + assert.Len(rArgs, 3) + + // ensure OR sub-filter is generated correctly + f = &filterBuilder{} + f.addWhere(clause1) + f.addWhere(clause2, arg1, arg2) + f.or(sf) + + r, rArgs = f.generateWhereClauses() + assert.Equal(fmt.Sprintf("((%s) AND (%s)) OR ((%s))", clause1, clause2, clause3), r) + assert.Len(rArgs, 3) + + // ensure NOT sub-filter is generated correctly + f = &filterBuilder{} + f.addWhere(clause1) + f.addWhere(clause2, arg1, arg2) + f.not(sf) + + r, rArgs = f.generateWhereClauses() + assert.Equal(fmt.Sprintf("((%s) AND (%s)) AND NOT ((%s))", clause1, clause2, clause3), r) + assert.Len(rArgs, 3) + + // ensure empty filter with ANDed sub-filter does not include AND + f = &filterBuilder{} + f.and(sf) + + r, rArgs = f.generateWhereClauses() + assert.Equal(fmt.Sprintf("((%s))", clause3), r) + assert.Len(rArgs, 1) + + // ensure empty filter with ORed sub-filter does not include OR + f = &filterBuilder{} + f.or(sf) + + r, rArgs = f.generateWhereClauses() + assert.Equal(fmt.Sprintf("((%s))", clause3), r) + assert.Len(rArgs, 1) + + // ensure empty filter with NOTed sub-filter does not include AND + f = &filterBuilder{} + f.not(sf) + + r, rArgs = f.generateWhereClauses() + assert.Equal(fmt.Sprintf("NOT ((%s))", clause3), r) + assert.Len(rArgs, 1) + + // (clause1) AND ((clause2) OR (clause3)) + f = &filterBuilder{} + f.addWhere(clause1) + sf2 := &filterBuilder{} + sf2.addWhere(clause2, arg1, arg2) + f.and(sf2) + sf2.or(sf) + r, rArgs = f.generateWhereClauses() + assert.Equal(fmt.Sprintf("(%s) AND ((%s) OR ((%s)))", clause1, clause2, clause3), r) + assert.Len(rArgs, 3) +} + +func TestGenerateHavingClauses(t *testing.T) { + assert := assert.New(t) + + f := &filterBuilder{} + + const clause1 = "a = 1" + const clause2 = "b = 2" + const clause3 = "c = 3" + + const arg1 = "1" + const arg2 = "2" + const arg3 = "3" + + // ensure single Having clause is generated correctly + f.addHaving(clause1) + r, rArgs := f.generateHavingClauses() + assert.Equal(fmt.Sprintf("(%s)", clause1), r) + assert.Len(rArgs, 0) + + // ensure multiple Having clauses are surrounded with parenthesis and + // ANDed together + f.addHaving(clause2, arg1, arg2) + r, rArgs = f.generateHavingClauses() + assert.Equal("(("+clause1+") AND ("+clause2+"))", r) + assert.Len(rArgs, 2) + + // ensure empty subfilter is not added to generated Having clause + sf := &filterBuilder{} + f.and(sf) + + r, rArgs = f.generateHavingClauses() + assert.Equal("(("+clause1+") AND ("+clause2+"))", r) + assert.Len(rArgs, 2) + + // ensure sub-filter is generated correctly + sf.addHaving(clause3, arg3) + r, rArgs = f.generateHavingClauses() + assert.Equal("(("+clause1+") AND ("+clause2+")) AND (("+clause3+"))", r) + assert.Len(rArgs, 3) + + // ensure OR sub-filter is generated correctly + f = &filterBuilder{} + f.addHaving(clause1) + f.addHaving(clause2, arg1, arg2) + f.or(sf) + + r, rArgs = f.generateHavingClauses() + assert.Equal("(("+clause1+") AND ("+clause2+")) OR (("+clause3+"))", r) + assert.Len(rArgs, 3) + + // ensure NOT sub-filter is generated correctly + f = &filterBuilder{} + f.addHaving(clause1) + f.addHaving(clause2, arg1, arg2) + f.not(sf) + + r, rArgs = f.generateHavingClauses() + assert.Equal("(("+clause1+") AND ("+clause2+")) AND NOT (("+clause3+"))", r) + assert.Len(rArgs, 3) +} + +func TestGetAllJoins(t *testing.T) { + assert := assert.New(t) + f := &filterBuilder{} + + const ( + table1Name = "table1Name" + table2Name = "table2Name" + + as1Name = "as1" + as2Name = "as2" + + onClause = "onClause1" + ) + + f.addLeftJoin(table1Name, as1Name, onClause) + + // ensure join is returned + joins := f.getAllJoins() + assert.Len(joins, 1) + assert.Equal(fmt.Sprintf("LEFT JOIN %s AS %s ON %s", table1Name, as1Name, onClause), joins[0].toSQL()) + + // ensure joins in sub-filter are returned + subFilter := &filterBuilder{} + f.and(subFilter) + subFilter.addLeftJoin(table2Name, as2Name, onClause) + + joins = f.getAllJoins() + assert.Len(joins, 2) + assert.Equal(fmt.Sprintf("LEFT JOIN %s AS %s ON %s", table2Name, as2Name, onClause), joins[1].toSQL()) + + // ensure redundant joins are not returned + subFilter.addLeftJoin(as1Name, "", onClause) + joins = f.getAllJoins() + assert.Len(joins, 2) +} + +func TestGetError(t *testing.T) { + assert := assert.New(t) + f := &filterBuilder{} + subFilter := &filterBuilder{} + + f.and(subFilter) + + expectedErr := errors.New("test error") + expectedErr2 := errors.New("test error2") + f.err = expectedErr + subFilter.err = expectedErr2 + + // ensure getError returns the top-level error state + assert.Equal(expectedErr, f.getError()) + + // ensure getError returns sub-filter error state if top-level error + // is nil + f.err = nil + assert.Equal(expectedErr2, f.getError()) + + // ensure getError returns nil if all error states are nil + subFilter.err = nil + assert.Nil(f.getError()) +} + +func TestStringCriterionHandlerIncludes(t *testing.T) { + assert := assert.New(t) + + const column = "column" + const value1 = "two words" + const quotedValue = `"two words"` + + f := &filterBuilder{} + f.handleCriterion(testCtx, stringCriterionHandler(&models.StringCriterionInput{ + Modifier: models.CriterionModifierIncludes, + Value: value1, + }, column)) + + assert.Len(f.whereClauses, 1) + assert.Equal(fmt.Sprintf("(%[1]s ILIKE ? OR %[1]s ILIKE ?)", column), f.whereClauses[0].sql) + assert.Len(f.whereClauses[0].args, 2) + assert.Equal("%two%", f.whereClauses[0].args[0]) + assert.Equal("%words%", f.whereClauses[0].args[1]) + + f = &filterBuilder{} + f.handleCriterion(testCtx, stringCriterionHandler(&models.StringCriterionInput{ + Modifier: models.CriterionModifierIncludes, + Value: quotedValue, + }, column)) + + assert.Len(f.whereClauses, 1) + assert.Equal(fmt.Sprintf("(%[1]s ILIKE ?)", column), f.whereClauses[0].sql) + assert.Len(f.whereClauses[0].args, 1) + assert.Equal("%two words%", f.whereClauses[0].args[0]) +} + +func TestStringCriterionHandlerExcludes(t *testing.T) { + assert := assert.New(t) + + const column = "column" + const value1 = "two words" + const quotedValue = `"two words"` + + f := &filterBuilder{} + f.handleCriterion(testCtx, stringCriterionHandler(&models.StringCriterionInput{ + Modifier: models.CriterionModifierExcludes, + Value: value1, + }, column)) + + assert.Len(f.whereClauses, 1) + assert.Equal(fmt.Sprintf("(%[1]s NOT ILIKE ? AND %[1]s NOT ILIKE ?)", column), f.whereClauses[0].sql) + assert.Len(f.whereClauses[0].args, 2) + assert.Equal("%two%", f.whereClauses[0].args[0]) + assert.Equal("%words%", f.whereClauses[0].args[1]) + + f = &filterBuilder{} + f.handleCriterion(testCtx, stringCriterionHandler(&models.StringCriterionInput{ + Modifier: models.CriterionModifierExcludes, + Value: quotedValue, + }, column)) + + assert.Len(f.whereClauses, 1) + assert.Equal(fmt.Sprintf("(%[1]s NOT ILIKE ?)", column), f.whereClauses[0].sql) + assert.Len(f.whereClauses[0].args, 1) + assert.Equal("%two words%", f.whereClauses[0].args[0]) +} + +func TestStringCriterionHandlerEquals(t *testing.T) { + assert := assert.New(t) + + const column = "column" + const value1 = "two words" + + f := &filterBuilder{} + f.handleCriterion(testCtx, stringCriterionHandler(&models.StringCriterionInput{ + Modifier: models.CriterionModifierEquals, + Value: value1, + }, column)) + + assert.Len(f.whereClauses, 1) + assert.Equal(fmt.Sprintf("%[1]s ILIKE ?", column), f.whereClauses[0].sql) + assert.Len(f.whereClauses[0].args, 1) + assert.Equal(value1, f.whereClauses[0].args[0]) +} + +func TestStringCriterionHandlerNotEquals(t *testing.T) { + assert := assert.New(t) + + const column = "column" + const value1 = "two words" + + f := &filterBuilder{} + f.handleCriterion(testCtx, stringCriterionHandler(&models.StringCriterionInput{ + Modifier: models.CriterionModifierNotEquals, + Value: value1, + }, column)) + + assert.Len(f.whereClauses, 1) + assert.Equal(fmt.Sprintf("%[1]s NOT ILIKE ?", column), f.whereClauses[0].sql) + assert.Len(f.whereClauses[0].args, 1) + assert.Equal(value1, f.whereClauses[0].args[0]) +} + +func TestStringCriterionHandlerMatchesRegex(t *testing.T) { + assert := assert.New(t) + + const column = "column" + const validValue = "two words" + const invalidValue = "*two words" + + f := &filterBuilder{} + f.handleCriterion(testCtx, stringCriterionHandler(&models.StringCriterionInput{ + Modifier: models.CriterionModifierMatchesRegex, + Value: validValue, + }, column)) + + assert.Len(f.whereClauses, 1) + assert.Equal(fmt.Sprintf("(%s IS NOT NULL AND regex_match(%[1]s, ?))", column), f.whereClauses[0].sql) + assert.Len(f.whereClauses[0].args, 1) + assert.Equal(validValue, f.whereClauses[0].args[0]) + + // ensure invalid regex sets error state + f = &filterBuilder{} + f.handleCriterion(testCtx, stringCriterionHandler(&models.StringCriterionInput{ + Modifier: models.CriterionModifierMatchesRegex, + Value: invalidValue, + }, column)) + + assert.NotNil(f.getError()) +} + +func TestStringCriterionHandlerNotMatchesRegex(t *testing.T) { + assert := assert.New(t) + + const column = "column" + const validValue = "two words" + const invalidValue = "*two words" + + f := &filterBuilder{} + f.handleCriterion(testCtx, stringCriterionHandler(&models.StringCriterionInput{ + Modifier: models.CriterionModifierNotMatchesRegex, + Value: validValue, + }, column)) + + assert.Len(f.whereClauses, 1) + assert.Equal(fmt.Sprintf("(%s IS NULL OR NOT regex_match(%[1]s, ?))", column), f.whereClauses[0].sql) + assert.Len(f.whereClauses[0].args, 1) + assert.Equal(validValue, f.whereClauses[0].args[0]) + + // ensure invalid regex sets error state + f = &filterBuilder{} + f.handleCriterion(testCtx, stringCriterionHandler(&models.StringCriterionInput{ + Modifier: models.CriterionModifierNotMatchesRegex, + Value: invalidValue, + }, column)) + + assert.NotNil(f.getError()) +} + +func TestStringCriterionHandlerIsNull(t *testing.T) { + assert := assert.New(t) + + const column = "column" + + f := &filterBuilder{} + f.handleCriterion(testCtx, stringCriterionHandler(&models.StringCriterionInput{ + Modifier: models.CriterionModifierIsNull, + }, column)) + + assert.Len(f.whereClauses, 1) + assert.Equal(fmt.Sprintf("(%[1]s IS NULL OR TRIM(%[1]s) = '')", column), f.whereClauses[0].sql) + assert.Len(f.whereClauses[0].args, 0) +} + +func TestStringCriterionHandlerNotNull(t *testing.T) { + assert := assert.New(t) + + const column = "column" + + f := &filterBuilder{} + f.handleCriterion(testCtx, stringCriterionHandler(&models.StringCriterionInput{ + Modifier: models.CriterionModifierNotNull, + }, column)) + + assert.Len(f.whereClauses, 1) + assert.Equal(fmt.Sprintf("(%[1]s IS NOT NULL AND TRIM(%[1]s) != '')", column), f.whereClauses[0].sql) + assert.Len(f.whereClauses[0].args, 0) +} diff --git a/pkg/postgres/fingerprint.go b/pkg/postgres/fingerprint.go new file mode 100644 index 0000000000..9a49c33bb9 --- /dev/null +++ b/pkg/postgres/fingerprint.go @@ -0,0 +1,129 @@ +package postgres + +import ( + "context" + "fmt" + "strconv" + "strings" + + "github.com/doug-martin/goqu/v9" + "github.com/doug-martin/goqu/v9/exp" + "github.com/stashapp/stash/pkg/models" + "gopkg.in/guregu/null.v4" +) + +const ( + fingerprintTable = "files_fingerprints" +) + +type fingerprintQueryRow struct { + Type null.String `db:"fingerprint_type"` + Fingerprint interface{} `db:"fingerprint"` +} + +func (r fingerprintQueryRow) valid() bool { + return r.Type.Valid +} + +func (r *fingerprintQueryRow) correct() { + if !r.Type.Valid || strings.ToLower(r.Type.String) != "phash" { + return + } + + if val, ok := r.Fingerprint.(string); ok { + if i, err := strconv.ParseInt(val, 10, 64); err == nil { + r.Fingerprint = i + } + } +} + +func (r *fingerprintQueryRow) resolve() models.Fingerprint { + return models.Fingerprint{ + Type: r.Type.String, + Fingerprint: r.Fingerprint, + } +} + +type fingerprintQueryBuilder struct { + repository + + tableMgr *table +} + +var FingerprintReaderWriter = &fingerprintQueryBuilder{ + repository: repository{ + tableName: fingerprintTable, + idColumn: fileIDColumn, + }, + + tableMgr: fingerprintTableMgr, +} + +func (qb *fingerprintQueryBuilder) insert(ctx context.Context, fileID models.FileID, f models.Fingerprint) error { + table := qb.table() + q := dialect.Insert(table).Cols(fileIDColumn, "type", "fingerprint").Vals( + goqu.Vals{fileID, f.Type, f.Fingerprint}, + ) + _, err := exec(ctx, q) + if err != nil { + return fmt.Errorf("inserting into %s: %w", table.GetTable(), err) + } + + return nil +} + +func (qb *fingerprintQueryBuilder) insertJoins(ctx context.Context, fileID models.FileID, f []models.Fingerprint) error { + for _, ff := range f { + if err := qb.insert(ctx, fileID, ff); err != nil { + return err + } + } + + return nil +} + +func (qb *fingerprintQueryBuilder) upsertJoins(ctx context.Context, fileID models.FileID, f []models.Fingerprint) error { + types := make([]string, len(f)) + for i, ff := range f { + types[i] = ff.Type + } + + if err := qb.destroyJoins(ctx, fileID, types); err != nil { + return err + } + + for _, ff := range f { + if err := qb.insert(ctx, fileID, ff); err != nil { + return err + } + } + + return nil +} + +func (qb *fingerprintQueryBuilder) replaceJoins(ctx context.Context, fileID models.FileID, f []models.Fingerprint) error { + if err := qb.destroy(ctx, []int{int(fileID)}); err != nil { + return err + } + + return qb.insertJoins(ctx, fileID, f) +} + +func (qb *fingerprintQueryBuilder) destroyJoins(ctx context.Context, fileID models.FileID, types []string) error { + table := qb.table() + q := dialect.Delete(table).Where( + table.Col(fileIDColumn).Eq(fileID), + table.Col("type").In(types), + ) + + _, err := exec(ctx, q) + if err != nil { + return fmt.Errorf("deleting from %s: %w", table.GetTable(), err) + } + + return nil +} + +func (qb *fingerprintQueryBuilder) table() exp.IdentifierExpression { + return qb.tableMgr.table +} diff --git a/pkg/postgres/folder.go b/pkg/postgres/folder.go new file mode 100644 index 0000000000..3e7e6c5db3 --- /dev/null +++ b/pkg/postgres/folder.go @@ -0,0 +1,551 @@ +package postgres + +import ( + "context" + "database/sql" + "errors" + "fmt" + "path/filepath" + "slices" + + "github.com/doug-martin/goqu/v9" + "github.com/doug-martin/goqu/v9/exp" + "github.com/jmoiron/sqlx" + "github.com/stashapp/stash/pkg/models" + "gopkg.in/guregu/null.v4" +) + +const folderTable = "folders" +const folderIDColumn = "folder_id" + +type folderRow struct { + ID models.FolderID `db:"id" goqu:"skipinsert"` + Path string `db:"path"` + ZipFileID null.Int `db:"zip_file_id"` + ParentFolderID null.Int `db:"parent_folder_id"` + ModTime Timestamp `db:"mod_time"` + CreatedAt Timestamp `db:"created_at"` + UpdatedAt Timestamp `db:"updated_at"` +} + +func (r *folderRow) fromFolder(o models.Folder) { + r.ID = o.ID + r.Path = o.Path + r.ZipFileID = nullIntFromFileIDPtr(o.ZipFileID) + r.ParentFolderID = nullIntFromFolderIDPtr(o.ParentFolderID) + r.ModTime = Timestamp{Timestamp: o.ModTime} + r.CreatedAt = Timestamp{Timestamp: o.CreatedAt} + r.UpdatedAt = Timestamp{Timestamp: o.UpdatedAt} +} + +type folderQueryRow struct { + folderRow + + ZipBasename null.String `db:"zip_basename"` + ZipFolderPath null.String `db:"zip_folder_path"` + ZipSize null.Int `db:"zip_size"` +} + +func (r *folderQueryRow) resolve() *models.Folder { + ret := &models.Folder{ + ID: r.ID, + DirEntry: models.DirEntry{ + ZipFileID: nullIntFileIDPtr(r.ZipFileID), + ModTime: r.ModTime.Timestamp.UTC(), + }, + Path: string(r.Path), + ParentFolderID: nullIntFolderIDPtr(r.ParentFolderID), + CreatedAt: r.CreatedAt.Timestamp.UTC(), + UpdatedAt: r.UpdatedAt.Timestamp.UTC(), + } + + if ret.ZipFileID != nil && r.ZipFolderPath.Valid && r.ZipBasename.Valid { + ret.ZipFile = &models.BaseFile{ + ID: *ret.ZipFileID, + Path: filepath.Join(r.ZipFolderPath.String, r.ZipBasename.String), + Basename: r.ZipBasename.String, + Size: r.ZipSize.Int64, + } + } + + return ret +} + +type folderQueryRows []folderQueryRow + +func (r folderQueryRows) resolve() []*models.Folder { + var ret []*models.Folder + + for _, row := range r { + f := row.resolve() + ret = append(ret, f) + } + + return ret +} + +type folderRepositoryType struct { + repository + + galleries repository +} + +var ( + folderRepository = folderRepositoryType{ + repository: repository{ + tableName: folderTable, + idColumn: idColumn, + }, + galleries: repository{ + tableName: galleryTable, + idColumn: folderIDColumn, + }, + } +) + +type FolderStore struct { + repository + + tableMgr *table +} + +func NewFolderStore() *FolderStore { + return &FolderStore{ + repository: repository{ + tableName: folderTable, + idColumn: idColumn, + }, + + tableMgr: folderTableMgr, + } +} + +func (qb *FolderStore) Create(ctx context.Context, f *models.Folder) error { + var r folderRow + r.fromFolder(*f) + + id, err := qb.tableMgr.insertID(ctx, r) + if err != nil { + return err + } + + // only assign id once we are successful + f.ID = models.FolderID(id) + + return nil +} + +func (qb *FolderStore) Update(ctx context.Context, updatedObject *models.Folder) error { + var r folderRow + r.fromFolder(*updatedObject) + + if err := qb.tableMgr.updateByID(ctx, updatedObject.ID, r); err != nil { + return err + } + + return nil +} + +func (qb *FolderStore) Destroy(ctx context.Context, id models.FolderID) error { + return qb.tableMgr.destroyExisting(ctx, []int{int(id)}) +} + +func (qb *FolderStore) table() exp.IdentifierExpression { + return qb.tableMgr.table +} + +func (qb *FolderStore) selectDataset() *goqu.SelectDataset { + table := qb.table() + fileTable := fileTableMgr.table + + zipFileTable := fileTable.As("zip_files") + zipFolderTable := table.As("zip_files_folders") + + cols := []interface{}{ + table.Col("id"), + table.Col("path"), + table.Col("zip_file_id"), + table.Col("parent_folder_id"), + table.Col("mod_time"), + table.Col("created_at"), + table.Col("updated_at"), + zipFileTable.Col("basename").As("zip_basename"), + zipFolderTable.Col("path").As("zip_folder_path"), + // size is needed to open containing zip files + zipFileTable.Col("size").As("zip_size"), + } + + ret := dialect.From(table).Select(cols...) + + return ret.LeftJoin( + zipFileTable, + goqu.On(table.Col("zip_file_id").Eq(zipFileTable.Col("id"))), + ).LeftJoin( + zipFolderTable, + goqu.On(zipFileTable.Col("parent_folder_id").Eq(zipFolderTable.Col(idColumn))), + ) +} + +func (qb *FolderStore) countDataset() *goqu.SelectDataset { + table := qb.table() + fileTable := fileTableMgr.table + + zipFileTable := fileTable.As("zip_files") + zipFolderTable := table.As("zip_files_folders") + + ret := dialect.From(table).Select(goqu.COUNT(goqu.DISTINCT(table.Col("id")))) + + return ret.LeftJoin( + zipFileTable, + goqu.On(table.Col("zip_file_id").Eq(zipFileTable.Col("id"))), + ).LeftJoin( + zipFolderTable, + goqu.On(zipFileTable.Col("parent_folder_id").Eq(zipFolderTable.Col(idColumn))), + ) +} + +func (qb *FolderStore) get(ctx context.Context, q *goqu.SelectDataset) (*models.Folder, error) { + ret, err := qb.getMany(ctx, q) + if err != nil { + return nil, err + } + + if len(ret) == 0 { + return nil, sql.ErrNoRows + } + + return ret[0], nil +} + +func (qb *FolderStore) getMany(ctx context.Context, q *goqu.SelectDataset) ([]*models.Folder, error) { + const single = false + var rows folderQueryRows + if err := queryFunc(ctx, q, single, func(r *sqlx.Rows) error { + var f folderQueryRow + if err := r.StructScan(&f); err != nil { + return err + } + + rows = append(rows, f) + return nil + }); err != nil { + return nil, err + } + + return rows.resolve(), nil +} + +func (qb *FolderStore) Find(ctx context.Context, id models.FolderID) (*models.Folder, error) { + q := qb.selectDataset().Where(qb.tableMgr.byID(id)) + + ret, err := qb.get(ctx, q) + if err != nil { + return nil, fmt.Errorf("getting folder by id %d: %w", id, err) + } + + return ret, nil +} + +// FindByIDs finds multiple folders by their IDs. +// No check is made to see if the folders exist, and the order of the returned folders +// is not guaranteed to be the same as the order of the input IDs. +func (qb *FolderStore) FindByIDs(ctx context.Context, ids []models.FolderID) ([]*models.Folder, error) { + folders := make([]*models.Folder, 0, len(ids)) + + table := qb.table() + if err := batchExec(ids, defaultBatchSize, func(batch []models.FolderID) error { + q := qb.selectDataset().Prepared(true).Where(table.Col(idColumn).In(batch)) + unsorted, err := qb.getMany(ctx, q) + if err != nil { + return err + } + + folders = append(folders, unsorted...) + + return nil + }); err != nil { + return nil, err + } + + return folders, nil +} + +func (qb *FolderStore) FindMany(ctx context.Context, ids []models.FolderID) ([]*models.Folder, error) { + folders := make([]*models.Folder, len(ids)) + + unsorted, err := qb.FindByIDs(ctx, ids) + if err != nil { + return nil, err + } + + for _, s := range unsorted { + i := slices.Index(ids, s.ID) + folders[i] = s + } + + for i := range folders { + if folders[i] == nil { + return nil, fmt.Errorf("folder with id %d not found", ids[i]) + } + } + + return folders, nil +} + +func (qb *FolderStore) FindByPath(ctx context.Context, p string, caseSensitive bool) (*models.Folder, error) { + // use like for case insensitive search + var criterion exp.BooleanExpression + if caseSensitive { + criterion = qb.table().Col("path").Eq(p) + } else { + criterion = qb.table().Col("path").ILike(p) + } + + q := qb.selectDataset().Prepared(true).Where(criterion) + + ret, err := qb.get(ctx, q) + if err != nil && !errors.Is(err, sql.ErrNoRows) { + return nil, fmt.Errorf("getting folder by path %s: %w", p, err) + } + + return ret, nil +} + +func (qb *FolderStore) FindByParentFolderID(ctx context.Context, parentFolderID models.FolderID) ([]*models.Folder, error) { + q := qb.selectDataset().Where(qb.table().Col("parent_folder_id").Eq(int(parentFolderID))) + + ret, err := qb.getMany(ctx, q) + if err != nil { + return nil, fmt.Errorf("getting folders by parent folder id %d: %w", parentFolderID, err) + } + + return ret, nil +} + +func (qb *FolderStore) allInPaths(q *goqu.SelectDataset, p []string) *goqu.SelectDataset { + table := qb.table() + + var conds []exp.Expression + for _, pp := range p { + ppWildcard := pp + string(filepath.Separator) + "%" + + conds = append(conds, table.Col("path").Eq(pp), table.Col("path").ILike(ppWildcard)) + } + + return q.Where( + goqu.Or(conds...), + ) +} + +// FindAllInPaths returns the all folders that are or are within any of the given paths. +// Returns all if limit is < 0. +// Returns all folders if p is empty. +func (qb *FolderStore) FindAllInPaths(ctx context.Context, p []string, limit, offset int) ([]*models.Folder, error) { + q := qb.selectDataset().Prepared(true) + q = qb.allInPaths(q, p) + + if limit > -1 { + q = q.Limit(uint(limit)) + } + + q = q.Offset(uint(offset)) + + ret, err := qb.getMany(ctx, q) + if err != nil && !errors.Is(err, sql.ErrNoRows) { + return nil, fmt.Errorf("getting folders in path %s: %w", p, err) + } + + return ret, nil +} + +// CountAllInPaths returns a count of all folders that are within any of the given paths. +// Returns count of all folders if p is empty. +func (qb *FolderStore) CountAllInPaths(ctx context.Context, p []string) (int, error) { + q := qb.countDataset().Prepared(true) + q = qb.allInPaths(q, p) + + return count(ctx, q) +} + +// func (qb *FolderStore) findBySubquery(ctx context.Context, sq *goqu.SelectDataset) ([]*file.Folder, error) { +// table := qb.table() + +// q := qb.selectDataset().Prepared(true).Where( +// table.Col(idColumn).Eq( +// sq, +// ), +// ) + +// return qb.getMany(ctx, q) +// } + +func (qb *FolderStore) FindByZipFileID(ctx context.Context, zipFileID models.FileID) ([]*models.Folder, error) { + table := qb.table() + + q := qb.selectDataset().Prepared(true).Where( + table.Col("zip_file_id").Eq(zipFileID), + ) + + return qb.getMany(ctx, q) +} + +func (qb *FolderStore) validateFilter(fileFilter *models.FolderFilterType) error { + const and = "AND" + const or = "OR" + const not = "NOT" + + if fileFilter.And != nil { + if fileFilter.Or != nil { + return illegalFilterCombination(and, or) + } + if fileFilter.Not != nil { + return illegalFilterCombination(and, not) + } + + return qb.validateFilter(fileFilter.And) + } + + if fileFilter.Or != nil { + if fileFilter.Not != nil { + return illegalFilterCombination(or, not) + } + + return qb.validateFilter(fileFilter.Or) + } + + if fileFilter.Not != nil { + return qb.validateFilter(fileFilter.Not) + } + + return nil +} + +func (qb *FolderStore) makeFilter(ctx context.Context, folderFilter *models.FolderFilterType) *filterBuilder { + query := &filterBuilder{} + + if folderFilter.And != nil { + query.and(qb.makeFilter(ctx, folderFilter.And)) + } + if folderFilter.Or != nil { + query.or(qb.makeFilter(ctx, folderFilter.Or)) + } + if folderFilter.Not != nil { + query.not(qb.makeFilter(ctx, folderFilter.Not)) + } + + filter := filterBuilderFromHandler(ctx, &folderFilterHandler{ + folderFilter: folderFilter, + }) + + return filter +} + +func (qb *FolderStore) Query(ctx context.Context, options models.FolderQueryOptions) (*models.FolderQueryResult, error) { + folderFilter := options.FolderFilter + findFilter := options.FindFilter + + if folderFilter == nil { + folderFilter = &models.FolderFilterType{} + } + if findFilter == nil { + findFilter = &models.FindFilterType{} + } + + query := qb.newQuery() + + distinctIDs(&query, folderTable) + + if q := findFilter.Q; q != nil && *q != "" { + searchColumns := []string{"folders.path"} + query.parseQueryString(searchColumns, *q) + } + + if err := qb.validateFilter(folderFilter); err != nil { + return nil, err + } + filter := qb.makeFilter(ctx, folderFilter) + + if err := query.addFilter(filter); err != nil { + return nil, err + } + + if err := qb.setQuerySort(&query, findFilter); err != nil { + return nil, err + } + query.pagination = getPagination(findFilter) + + result, err := qb.queryGroupedFields(ctx, options, query) + if err != nil { + return nil, fmt.Errorf("error querying aggregate fields: %w", err) + } + + idsResult, err := query.findIDs(ctx) + if err != nil { + return nil, fmt.Errorf("error finding IDs: %w", err) + } + + result.IDs = make([]models.FolderID, len(idsResult)) + for i, id := range idsResult { + result.IDs[i] = models.FolderID(id) + } + + return result, nil +} + +func (qb *FolderStore) queryGroupedFields(ctx context.Context, options models.FolderQueryOptions, query queryBuilder) (*models.FolderQueryResult, error) { + if !options.Count { + // nothing to do - return empty result + return models.NewFolderQueryResult(qb), nil + } + + aggregateQuery := qb.newQuery() + + if options.Count { + aggregateQuery.addColumn("COUNT(DISTINCT temp.id) as total") + } + + const includeSortPagination = false + aggregateQuery.from = fmt.Sprintf("(%s) as temp", query.toSQL(includeSortPagination)) + + out := struct { + Total int + Duration float64 + Megapixels float64 + Size int64 + }{} + if err := qb.repository.queryStruct(ctx, aggregateQuery.toSQL(includeSortPagination), query.args, &out); err != nil { + return nil, err + } + + ret := models.NewFolderQueryResult(qb) + ret.Count = out.Total + + return ret, nil +} + +var folderSortOptions = sortOptions{ + "created_at", + "id", + "path", + "random", + "updated_at", +} + +func (qb *FolderStore) setQuerySort(query *queryBuilder, findFilter *models.FindFilterType) error { + if findFilter == nil || findFilter.Sort == nil || *findFilter.Sort == "" { + return nil + } + sort := findFilter.GetSort("path") + + // CVE-2024-32231 - ensure sort is in the list of allowed sorts + if err := folderSortOptions.validateSort(sort); err != nil { + return err + } + + direction := findFilter.GetDirection() + var agg []string + query.sort, agg = getSort(sort, direction, "folders") + query.addGroupBy(agg...) + + return nil +} diff --git a/pkg/postgres/folder_filter.go b/pkg/postgres/folder_filter.go new file mode 100644 index 0000000000..9f5ac421bf --- /dev/null +++ b/pkg/postgres/folder_filter.go @@ -0,0 +1,160 @@ +package postgres + +import ( + "context" + "fmt" + + "github.com/stashapp/stash/pkg/models" +) + +type folderFilterHandler struct { + folderFilter *models.FolderFilterType + table sqlTable + isRelated bool +} + +func (qb *folderFilterHandler) validate() error { + folderFilter := qb.folderFilter + if folderFilter == nil { + return nil + } + + if err := validateFilterCombination(folderFilter.OperatorFilter); err != nil { + return err + } + + if qb.isRelated && (folderFilter.GalleriesFilter != nil) { + return fmt.Errorf("cannot use related filters inside a related filter") + } + + if subFilter := folderFilter.SubFilter(); subFilter != nil { + sqb := &folderFilterHandler{folderFilter: subFilter, isRelated: qb.isRelated} + if err := sqb.validate(); err != nil { + return err + } + } + + return nil +} + +func (qb *folderFilterHandler) handle(ctx context.Context, f *filterBuilder) { + folderFilter := qb.folderFilter + if folderFilter == nil { + return + } + + if err := qb.validate(); err != nil { + f.setError(err) + return + } + + sf := folderFilter.SubFilter() + if sf != nil { + sub := &folderFilterHandler{folderFilter: sf, table: qb.table} + handleSubFilter(ctx, sub, f, folderFilter.OperatorFilter) + } + + f.handleCriterion(ctx, qb.criterionHandler()) +} + +func (qb *folderFilterHandler) criterionHandler() criterionHandler { + if qb.table == "" { + qb.table = folderTable + } + + folderFilter := qb.folderFilter + return compoundHandler{ + stringCriterionHandler(folderFilter.Path, qb.table.Col("path")), + ×tampCriterionHandler{folderFilter.ModTime, qb.table.Col("mod_time"), nil}, + + qb.parentFolderCriterionHandler(folderFilter.ParentFolder), + qb.zipFileCriterionHandler(folderFilter.ZipFile), + + qb.galleryCountCriterionHandler(folderFilter.GalleryCount), + + ×tampCriterionHandler{folderFilter.CreatedAt, qb.table.Col("created_at"), nil}, + ×tampCriterionHandler{folderFilter.UpdatedAt, qb.table.Col("updated_at"), nil}, + + &relatedFilterHandler{ + relatedIDCol: qb.table.Col("id"), + relatedRepo: galleryRepository.repository, + relatedHandler: &galleryFilterHandler{folderFilter.GalleriesFilter}, + joinFn: func(f *filterBuilder) { + folderRepository.galleries.innerJoin(f, "", qb.table.Col("id")) + }, + }, + } +} + +func (qb *folderFilterHandler) zipFileCriterionHandler(criterion *models.MultiCriterionInput) criterionHandlerFunc { + return func(ctx context.Context, f *filterBuilder) { + if criterion != nil { + if criterion.Modifier == models.CriterionModifierIsNull || criterion.Modifier == models.CriterionModifierNotNull { + var notClause string + if criterion.Modifier == models.CriterionModifierNotNull { + notClause = "NOT" + } + + f.addWhere(fmt.Sprintf("%s.zip_file_id IS %s NULL", qb.table.Name(), notClause)) + return + } + + if len(criterion.Value) == 0 { + return + } + + var args []interface{} + for _, tagID := range criterion.Value { + args = append(args, tagID) + } + + whereClause := "" + havingClause := "" + switch criterion.Modifier { + case models.CriterionModifierIncludes: + whereClause = fmt.Sprintf("%s.zip_file_id IN %s", qb.table.Name(), getInBinding(len(criterion.Value))) + case models.CriterionModifierExcludes: + whereClause = fmt.Sprintf("%s.zip_file_id NOT IN %s", qb.table.Name(), getInBinding(len(criterion.Value))) + } + + f.addWhere(whereClause, args...) + f.addHaving(havingClause) + } + } +} + +func (qb *folderFilterHandler) parentFolderCriterionHandler(folder *models.HierarchicalMultiCriterionInput) criterionHandlerFunc { + return func(ctx context.Context, f *filterBuilder) { + if folder == nil { + return + } + + folderCopy := *folder + switch folderCopy.Modifier { + case models.CriterionModifierEquals: + folderCopy.Modifier = models.CriterionModifierIncludesAll + case models.CriterionModifierNotEquals: + folderCopy.Modifier = models.CriterionModifierExcludes + } + + hh := hierarchicalMultiCriterionHandlerBuilder{ + primaryTable: qb.table.Name(), + foreignTable: qb.table.Name(), + foreignFK: "parent_folder_id", + parentFK: "parent_folder_id", + } + + hh.handler(&folderCopy)(ctx, f) + } +} + +func (qb *folderFilterHandler) galleryCountCriterionHandler(galleryCount *models.IntCriterionInput) criterionHandlerFunc { + return func(ctx context.Context, f *filterBuilder) { + if galleryCount != nil { + f.addLeftJoin("galleries", "", "galleries.folder_id = folders.id") + clause, args := getIntCriterionWhereClause("count(distinct galleries.id)", *galleryCount) + + f.addHaving(clause, args...) + } + } +} diff --git a/pkg/postgres/gallery.go b/pkg/postgres/gallery.go new file mode 100644 index 0000000000..4ab7a05c34 --- /dev/null +++ b/pkg/postgres/gallery.go @@ -0,0 +1,919 @@ +package postgres + +import ( + "context" + "database/sql" + "errors" + "fmt" + "path/filepath" + "slices" + + "github.com/doug-martin/goqu/v9" + "github.com/doug-martin/goqu/v9/exp" + "github.com/jmoiron/sqlx" + "github.com/stashapp/stash/pkg/models" + "gopkg.in/guregu/null.v4" + "gopkg.in/guregu/null.v4/zero" +) + +const ( + galleryTable = "galleries" + + galleriesFilesTable = "galleries_files" + performersGalleriesTable = "performers_galleries" + galleriesTagsTable = "galleries_tags" + galleriesImagesTable = "galleries_images" + galleriesScenesTable = "scenes_galleries" + galleryIDColumn = "gallery_id" + galleriesURLsTable = "gallery_urls" + galleriesURLColumn = "url" +) + +type galleryRow struct { + ID int `db:"id" goqu:"skipinsert"` + Title zero.String `db:"title"` + Code zero.String `db:"code"` + Date NullDate `db:"date"` + DatePrecision null.Int `db:"date_precision"` + Details zero.String `db:"details"` + Photographer zero.String `db:"photographer"` + // expressed as 1-100 + Rating null.Int `db:"rating"` + Organized bool `db:"organized"` + StudioID null.Int `db:"studio_id,omitempty"` + FolderID null.Int `db:"folder_id,omitempty"` + CreatedAt Timestamp `db:"created_at"` + UpdatedAt Timestamp `db:"updated_at"` +} + +func (r *galleryRow) fromGallery(o models.Gallery) { + r.ID = o.ID + r.Title = zero.StringFrom(o.Title) + r.Code = zero.StringFrom(o.Code) + r.Date = NullDateFromDatePtr(o.Date) + r.DatePrecision = datePrecisionFromDatePtr(o.Date) + r.Details = zero.StringFrom(o.Details) + r.Photographer = zero.StringFrom(o.Photographer) + r.Rating = intFromPtr(o.Rating) + r.Organized = o.Organized + r.StudioID = intFromPtr(o.StudioID) + r.FolderID = nullIntFromFolderIDPtr(o.FolderID) + r.CreatedAt = Timestamp{Timestamp: o.CreatedAt} + r.UpdatedAt = Timestamp{Timestamp: o.UpdatedAt} +} + +type galleryQueryRow struct { + galleryRow + FolderPath zero.String `db:"folder_path"` + PrimaryFileID null.Int `db:"primary_file_id"` + PrimaryFileFolderPath zero.String `db:"primary_file_folder_path"` + PrimaryFileBasename zero.String `db:"primary_file_basename"` + PrimaryFileChecksum zero.String `db:"primary_file_checksum"` +} + +func (r *galleryQueryRow) resolve() *models.Gallery { + ret := &models.Gallery{ + ID: r.ID, + Title: r.Title.String, + Code: r.Code.String, + Date: r.Date.DatePtr(r.DatePrecision), + Details: r.Details.String, + Photographer: r.Photographer.String, + Rating: nullIntPtr(r.Rating), + Organized: r.Organized, + StudioID: nullIntPtr(r.StudioID), + FolderID: nullIntFolderIDPtr(r.FolderID), + PrimaryFileID: nullIntFileIDPtr(r.PrimaryFileID), + CreatedAt: r.CreatedAt.Timestamp.UTC(), + UpdatedAt: r.UpdatedAt.Timestamp.UTC(), + } + + if r.PrimaryFileFolderPath.Valid && r.PrimaryFileBasename.Valid { + ret.Path = filepath.Join(r.PrimaryFileFolderPath.String, r.PrimaryFileBasename.String) + } else if r.FolderPath.Valid { + ret.Path = r.FolderPath.String + } + + return ret +} + +type galleryRowRecord struct { + updateRecord +} + +func (r *galleryRowRecord) fromPartial(o models.GalleryPartial) { + r.setNullString("title", o.Title) + r.setNullString("code", o.Code) + r.setNullDate("date", "date_precision", o.Date) + r.setNullString("details", o.Details) + r.setNullString("photographer", o.Photographer) + r.setNullInt("rating", o.Rating) + r.setBool("organized", o.Organized) + r.setNullInt("studio_id", o.StudioID) + r.setTimestamp("created_at", o.CreatedAt) + r.setTimestamp("updated_at", o.UpdatedAt) +} + +type galleryRepositoryType struct { + repository + performers joinRepository + images joinRepository + tags joinRepository + scenes joinRepository + files filesRepository +} + +func (r *galleryRepositoryType) addGalleriesFilesTable(f *filterBuilder) { + f.addLeftJoin(galleriesFilesTable, "", "galleries_files.gallery_id = galleries.id AND galleries_files.\"primary\" = true") +} + +func (r *galleryRepositoryType) addFilesTable(f *filterBuilder) { + r.addGalleriesFilesTable(f) + f.addLeftJoin(fileTable, "", "galleries_files.file_id = files.id") +} + +func (r *galleryRepositoryType) addFoldersTable(f *filterBuilder) { + r.addFilesTable(f) + f.addLeftJoin(folderTable, "", "files.parent_folder_id = folders.id") +} + +var ( + galleryRepository = galleryRepositoryType{ + repository: repository{ + tableName: galleryTable, + idColumn: idColumn, + }, + performers: joinRepository{ + repository: repository{ + tableName: performersGalleriesTable, + idColumn: galleryIDColumn, + }, + fkColumn: "performer_id", + }, + tags: joinRepository{ + repository: repository{ + tableName: galleriesTagsTable, + idColumn: galleryIDColumn, + }, + fkColumn: "tag_id", + foreignTable: tagTable, + orderBy: tagTableSortSQL, + }, + images: joinRepository{ + repository: repository{ + tableName: galleriesImagesTable, + idColumn: galleryIDColumn, + }, + fkColumn: "image_id", + }, + scenes: joinRepository{ + repository: repository{ + tableName: galleriesScenesTable, + idColumn: galleryIDColumn, + }, + fkColumn: sceneIDColumn, + }, + files: filesRepository{ + repository: repository{ + tableName: galleriesFilesTable, + idColumn: galleryIDColumn, + }, + }, + } +) + +type GalleryStore struct { + tableMgr *table + + fileStore *FileStore + folderStore *FolderStore +} + +func NewGalleryStore(fileStore *FileStore, folderStore *FolderStore) *GalleryStore { + return &GalleryStore{ + tableMgr: galleryTableMgr, + fileStore: fileStore, + folderStore: folderStore, + } +} + +func (qb *GalleryStore) table() exp.IdentifierExpression { + return qb.tableMgr.table +} + +func (qb *GalleryStore) selectDataset() *goqu.SelectDataset { + table := qb.table() + files := fileTableMgr.table + folders := folderTableMgr.table + galleryFolder := folderTableMgr.table.As("gallery_folder") + + return dialect.From(table).LeftJoin( + galleriesFilesJoinTable, + goqu.On( + galleriesFilesJoinTable.Col(galleryIDColumn).Eq(table.Col(idColumn)), + galleriesFilesJoinTable.Col("primary").IsTrue(), + ), + ).LeftJoin( + files, + goqu.On(files.Col(idColumn).Eq(galleriesFilesJoinTable.Col(fileIDColumn))), + ).LeftJoin( + folders, + goqu.On(folders.Col(idColumn).Eq(files.Col("parent_folder_id"))), + ).LeftJoin( + galleryFolder, + goqu.On(galleryFolder.Col(idColumn).Eq(table.Col("folder_id"))), + ).Select( + qb.table().All(), + galleriesFilesJoinTable.Col(fileIDColumn).As("primary_file_id"), + folders.Col("path").As("primary_file_folder_path"), + files.Col("basename").As("primary_file_basename"), + galleryFolder.Col("path").As("folder_path"), + ) +} + +func (qb *GalleryStore) Create(ctx context.Context, newObject *models.Gallery, fileIDs []models.FileID) error { + var r galleryRow + r.fromGallery(*newObject) + + id, err := qb.tableMgr.insertID(ctx, r) + if err != nil { + return err + } + + if len(fileIDs) > 0 { + const firstPrimary = true + if err := galleriesFilesTableMgr.insertJoins(ctx, id, firstPrimary, fileIDs); err != nil { + return err + } + } + + if newObject.URLs.Loaded() { + const startPos = 0 + if err := galleriesURLsTableMgr.insertJoins(ctx, id, startPos, newObject.URLs.List()); err != nil { + return err + } + } + if newObject.PerformerIDs.Loaded() { + if err := galleriesPerformersTableMgr.insertJoins(ctx, id, newObject.PerformerIDs.List()); err != nil { + return err + } + } + if newObject.TagIDs.Loaded() { + if err := galleriesTagsTableMgr.insertJoins(ctx, id, newObject.TagIDs.List()); err != nil { + return err + } + } + if newObject.SceneIDs.Loaded() { + if err := galleriesScenesTableMgr.insertJoins(ctx, id, newObject.SceneIDs.List()); err != nil { + return err + } + } + + updated, err := qb.find(ctx, id) + if err != nil { + return fmt.Errorf("finding after create: %w", err) + } + + *newObject = *updated + + return nil +} + +func (qb *GalleryStore) Update(ctx context.Context, updatedObject *models.Gallery) error { + var r galleryRow + r.fromGallery(*updatedObject) + + if err := qb.tableMgr.updateByID(ctx, updatedObject.ID, r); err != nil { + return err + } + + if updatedObject.URLs.Loaded() { + if err := galleriesURLsTableMgr.replaceJoins(ctx, updatedObject.ID, updatedObject.URLs.List()); err != nil { + return err + } + } + if updatedObject.PerformerIDs.Loaded() { + if err := galleriesPerformersTableMgr.replaceJoins(ctx, updatedObject.ID, updatedObject.PerformerIDs.List()); err != nil { + return err + } + } + if updatedObject.TagIDs.Loaded() { + if err := galleriesTagsTableMgr.replaceJoins(ctx, updatedObject.ID, updatedObject.TagIDs.List()); err != nil { + return err + } + } + if updatedObject.SceneIDs.Loaded() { + if err := galleriesScenesTableMgr.replaceJoins(ctx, updatedObject.ID, updatedObject.SceneIDs.List()); err != nil { + return err + } + } + + if updatedObject.Files.Loaded() { + fileIDs := make([]models.FileID, len(updatedObject.Files.List())) + for i, f := range updatedObject.Files.List() { + fileIDs[i] = f.Base().ID + } + + if err := galleriesFilesTableMgr.replaceJoins(ctx, updatedObject.ID, fileIDs); err != nil { + return err + } + } + + return nil +} + +func (qb *GalleryStore) UpdatePartial(ctx context.Context, id int, partial models.GalleryPartial) (*models.Gallery, error) { + r := galleryRowRecord{ + updateRecord{ + Record: make(exp.Record), + }, + } + + r.fromPartial(partial) + + if len(r.Record) > 0 { + if err := qb.tableMgr.updateByID(ctx, id, r.Record); err != nil { + return nil, err + } + } + + if partial.URLs != nil { + if err := galleriesURLsTableMgr.modifyJoins(ctx, id, partial.URLs.Values, partial.URLs.Mode); err != nil { + return nil, err + } + } + if partial.PerformerIDs != nil { + if err := galleriesPerformersTableMgr.modifyJoins(ctx, id, partial.PerformerIDs.IDs, partial.PerformerIDs.Mode); err != nil { + return nil, err + } + } + if partial.TagIDs != nil { + if err := galleriesTagsTableMgr.modifyJoins(ctx, id, partial.TagIDs.IDs, partial.TagIDs.Mode); err != nil { + return nil, err + } + } + if partial.SceneIDs != nil { + if err := galleriesScenesTableMgr.modifyJoins(ctx, id, partial.SceneIDs.IDs, partial.SceneIDs.Mode); err != nil { + return nil, err + } + } + + if partial.PrimaryFileID != nil { + if err := galleriesFilesTableMgr.setPrimary(ctx, id, *partial.PrimaryFileID); err != nil { + return nil, err + } + } + + return qb.find(ctx, id) +} + +func (qb *GalleryStore) Destroy(ctx context.Context, id int) error { + return qb.tableMgr.destroyExisting(ctx, []int{id}) +} + +func (qb *GalleryStore) GetFiles(ctx context.Context, id int) ([]models.File, error) { + fileIDs, err := galleryRepository.files.get(ctx, id) + if err != nil { + return nil, err + } + + // use fileStore to load files + files, err := qb.fileStore.Find(ctx, fileIDs...) + if err != nil { + return nil, err + } + + ret := make([]models.File, len(files)) + copy(ret, files) + + return ret, nil +} + +func (qb *GalleryStore) GetManyFileIDs(ctx context.Context, ids []int) ([][]models.FileID, error) { + const primaryOnly = false + return galleryRepository.files.getMany(ctx, ids, primaryOnly) +} + +// returns nil, nil if not found +func (qb *GalleryStore) Find(ctx context.Context, id int) (*models.Gallery, error) { + ret, err := qb.find(ctx, id) + if errors.Is(err, sql.ErrNoRows) { + return nil, nil + } + return ret, err +} + +func (qb *GalleryStore) FindMany(ctx context.Context, ids []int) ([]*models.Gallery, error) { + galleries := make([]*models.Gallery, len(ids)) + + if len(ids) == 0 { + return galleries, nil + } + + if err := batchExec(ids, defaultBatchSize, func(batch []int) error { + q := qb.selectDataset().Prepared(true).Where(qb.table().Col(idColumn).In(batch)) + unsorted, err := qb.getMany(ctx, q) + if err != nil { + return err + } + + for _, s := range unsorted { + i := slices.Index(ids, s.ID) + galleries[i] = s + } + + return nil + }); err != nil { + return nil, err + } + + for i := range galleries { + if galleries[i] == nil { + return nil, fmt.Errorf("gallery with id %d not found", ids[i]) + } + } + + return galleries, nil +} + +// returns nil, sql.ErrNoRows if not found +func (qb *GalleryStore) find(ctx context.Context, id int) (*models.Gallery, error) { + q := qb.selectDataset().Where(qb.tableMgr.byID(id)) + + ret, err := qb.get(ctx, q) + if err != nil { + return nil, err + } + + return ret, nil +} + +func (qb *GalleryStore) findBySubquery(ctx context.Context, sq *goqu.SelectDataset) ([]*models.Gallery, error) { + table := qb.table() + + q := qb.selectDataset().Prepared(true).Where( + table.Col(idColumn).Eq( + sq, + ), + ) + + return qb.getMany(ctx, q) +} + +// returns nil, sql.ErrNoRows if not found +func (qb *GalleryStore) get(ctx context.Context, q *goqu.SelectDataset) (*models.Gallery, error) { + ret, err := qb.getMany(ctx, q) + if err != nil { + return nil, err + } + + if len(ret) == 0 { + return nil, sql.ErrNoRows + } + + return ret[0], nil +} + +func (qb *GalleryStore) getMany(ctx context.Context, q *goqu.SelectDataset) ([]*models.Gallery, error) { + const single = false + var ret []*models.Gallery + var lastID int + if err := queryFunc(ctx, q, single, func(r *sqlx.Rows) error { + var f galleryQueryRow + if err := r.StructScan(&f); err != nil { + return err + } + + s := f.resolve() + + if s.ID == lastID { + return fmt.Errorf("internal error: multiple rows returned for single gallery id %d", s.ID) + } + lastID = s.ID + + ret = append(ret, s) + return nil + }); err != nil { + return nil, err + } + + return ret, nil +} + +func (qb *GalleryStore) FindByFileID(ctx context.Context, fileID models.FileID) ([]*models.Gallery, error) { + sq := dialect.From(galleriesFilesJoinTable).Select(galleriesFilesJoinTable.Col(galleryIDColumn)).Where( + galleriesFilesJoinTable.Col(fileIDColumn).Eq(fileID), + ) + + ret, err := qb.findBySubquery(ctx, sq) + if err != nil { + return nil, fmt.Errorf("getting gallery by file id %d: %w", fileID, err) + } + + return ret, nil +} + +func (qb *GalleryStore) CountByFileID(ctx context.Context, fileID models.FileID) (int, error) { + joinTable := galleriesFilesJoinTable + + q := dialect.Select(goqu.COUNT("*")).From(joinTable).Where(joinTable.Col(fileIDColumn).Eq(fileID)) + return count(ctx, q) +} + +func (qb *GalleryStore) FindByFingerprints(ctx context.Context, fp []models.Fingerprint) ([]*models.Gallery, error) { + fingerprintTable := fingerprintTableMgr.table + + var ex []exp.Expression + + for _, v := range fp { + ex = append(ex, goqu.And( + fingerprintTable.Col("type").Eq(v.Type), + fingerprintTable.Col("fingerprint").Eq(v.Fingerprint), + )) + } + + sq := dialect.From(galleriesFilesJoinTable). + InnerJoin( + fingerprintTable, + goqu.On(fingerprintTable.Col(fileIDColumn).Eq(galleriesFilesJoinTable.Col(fileIDColumn))), + ). + Select(galleriesFilesJoinTable.Col(galleryIDColumn)).Where(goqu.Or(ex...)) + + ret, err := qb.findBySubquery(ctx, sq) + if err != nil { + return nil, fmt.Errorf("getting gallery by fingerprints: %w", err) + } + + return ret, nil +} + +func (qb *GalleryStore) FindByChecksum(ctx context.Context, checksum string) ([]*models.Gallery, error) { + return qb.FindByFingerprints(ctx, []models.Fingerprint{ + { + Type: models.FingerprintTypeMD5, + Fingerprint: checksum, + }, + }) +} + +func (qb *GalleryStore) FindByChecksums(ctx context.Context, checksums []string) ([]*models.Gallery, error) { + fingerprints := make([]models.Fingerprint, len(checksums)) + + for i, c := range checksums { + fingerprints[i] = models.Fingerprint{ + Type: models.FingerprintTypeMD5, + Fingerprint: c, + } + } + return qb.FindByFingerprints(ctx, fingerprints) +} + +func (qb *GalleryStore) FindByPath(ctx context.Context, p string) ([]*models.Gallery, error) { + table := qb.table() + filesTable := fileTableMgr.table + fileFoldersTable := folderTableMgr.table.As("file_folders") + foldersTable := folderTableMgr.table + + basename := filepath.Base(p) + dir := filepath.Dir(p) + + sq := dialect.From(table).LeftJoin( + galleriesFilesJoinTable, + goqu.On(galleriesFilesJoinTable.Col(galleryIDColumn).Eq(table.Col(idColumn))), + ).LeftJoin( + filesTable, + goqu.On(filesTable.Col(idColumn).Eq(galleriesFilesJoinTable.Col(fileIDColumn))), + ).LeftJoin( + fileFoldersTable, + goqu.On(fileFoldersTable.Col(idColumn).Eq(filesTable.Col("parent_folder_id"))), + ).LeftJoin( + foldersTable, + goqu.On(foldersTable.Col(idColumn).Eq(table.Col("folder_id"))), + ).Select(table.Col(idColumn)).Where( + goqu.Or( + goqu.And( + fileFoldersTable.Col("path").Eq(dir), + filesTable.Col("basename").Eq(basename), + ), + foldersTable.Col("path").Eq(p), + ), + ) + + ret, err := qb.findBySubquery(ctx, sq) + if err != nil && !errors.Is(err, sql.ErrNoRows) { + return nil, fmt.Errorf("getting gallery by path %s: %w", p, err) + } + + return ret, nil +} + +func (qb *GalleryStore) FindByFolderID(ctx context.Context, folderID models.FolderID) ([]*models.Gallery, error) { + table := qb.table() + + sq := dialect.From(table).Select(table.Col(idColumn)).Where( + table.Col("folder_id").Eq(folderID), + ) + + ret, err := qb.findBySubquery(ctx, sq) + if err != nil { + return nil, fmt.Errorf("getting galleries for folder %d: %w", folderID, err) + } + + return ret, nil +} + +func (qb *GalleryStore) FindBySceneID(ctx context.Context, sceneID int) ([]*models.Gallery, error) { + sq := dialect.From(galleriesScenesJoinTable).Select(galleriesScenesJoinTable.Col(galleryIDColumn)).Where( + galleriesScenesJoinTable.Col(sceneIDColumn).Eq(sceneID), + ) + + ret, err := qb.findBySubquery(ctx, sq) + if err != nil { + return nil, fmt.Errorf("getting galleries for scene %d: %w", sceneID, err) + } + + return ret, nil +} + +func (qb *GalleryStore) FindByImageID(ctx context.Context, imageID int) ([]*models.Gallery, error) { + sq := dialect.From(galleriesImagesJoinTable).Select(galleriesImagesJoinTable.Col(galleryIDColumn)).Where( + galleriesImagesJoinTable.Col(imageIDColumn).Eq(imageID), + ) + + ret, err := qb.findBySubquery(ctx, sq) + if err != nil { + return nil, fmt.Errorf("getting galleries for image %d: %w", imageID, err) + } + + return ret, nil +} + +func (qb *GalleryStore) CountByImageID(ctx context.Context, imageID int) (int, error) { + joinTable := galleriesImagesJoinTable + + q := dialect.Select(goqu.COUNT("*")).From(joinTable).Where(joinTable.Col(imageIDColumn).Eq(imageID)) + return count(ctx, q) +} + +func (qb *GalleryStore) FindUserGalleryByTitle(ctx context.Context, title string) ([]*models.Gallery, error) { + table := qb.table() + + sq := dialect.From(table).LeftJoin( + galleriesFilesJoinTable, + goqu.On(galleriesFilesJoinTable.Col(galleryIDColumn).Eq(table.Col(idColumn))), + ).Select(table.Col(idColumn)).Where( + table.Col("folder_id").IsNull(), + galleriesFilesJoinTable.Col("file_id").IsNull(), + table.Col("title").Eq(title), + ) + + ret, err := qb.findBySubquery(ctx, sq) + if err != nil { + return nil, fmt.Errorf("getting user galleries for title %s: %w", title, err) + } + + return ret, nil +} + +func (qb *GalleryStore) Count(ctx context.Context) (int, error) { + q := dialect.Select(goqu.COUNT("*")).From(qb.table()) + return count(ctx, q) +} + +func (qb *GalleryStore) All(ctx context.Context) ([]*models.Gallery, error) { + return qb.getMany(ctx, qb.selectDataset()) +} + +func (qb *GalleryStore) makeQuery(ctx context.Context, galleryFilter *models.GalleryFilterType, findFilter *models.FindFilterType) (*queryBuilder, error) { + if galleryFilter == nil { + galleryFilter = &models.GalleryFilterType{} + } + if findFilter == nil { + findFilter = &models.FindFilterType{} + } + + query := galleryRepository.newQuery() + distinctIDs(&query, galleryTable) + + if q := findFilter.Q; q != nil && *q != "" { + query.addJoins( + join{ + table: galleriesFilesTable, + onClause: "galleries_files.gallery_id = galleries.id AND galleries_files.\"primary\" = true", + }, + join{ + table: fileTable, + onClause: "galleries_files.file_id = files.id", + }, + join{ + table: folderTable, + onClause: "files.parent_folder_id = folders.id", + }, + join{ + table: fingerprintTable, + onClause: "files_fingerprints.file_id = galleries_files.file_id", + }, + join{ + table: folderTable, + as: "gallery_folder", + onClause: "galleries.folder_id = gallery_folder.id", + }, + join{ + table: galleriesChaptersTable, + onClause: "galleries_chapters.gallery_id = galleries.id", + }, + ) + + // add joins for files and checksum + filepathColumn := "folders.path || '" + string(filepath.Separator) + "' || files.basename" + searchColumns := []string{"galleries.title", "gallery_folder.path", filepathColumn, "files_fingerprints.fingerprint", "galleries_chapters.title"} + query.parseQueryString(searchColumns, *q) + } + + filter := filterBuilderFromHandler(ctx, &galleryFilterHandler{ + galleryFilter: galleryFilter, + }) + + if err := query.addFilter(filter); err != nil { + return nil, err + } + + if err := qb.setGallerySort(&query, findFilter); err != nil { + return nil, err + } + query.pagination = getPagination(findFilter) + + return &query, nil +} + +func (qb *GalleryStore) Query(ctx context.Context, galleryFilter *models.GalleryFilterType, findFilter *models.FindFilterType) ([]*models.Gallery, int, error) { + query, err := qb.makeQuery(ctx, galleryFilter, findFilter) + if err != nil { + return nil, 0, err + } + + idsResult, countResult, err := query.executeFind(ctx) + if err != nil { + return nil, 0, err + } + + galleries, err := qb.FindMany(ctx, idsResult) + if err != nil { + return nil, 0, err + } + + return galleries, countResult, nil +} + +func (qb *GalleryStore) QueryCount(ctx context.Context, galleryFilter *models.GalleryFilterType, findFilter *models.FindFilterType) (int, error) { + query, err := qb.makeQuery(ctx, galleryFilter, findFilter) + if err != nil { + return 0, err + } + + return query.executeCount(ctx) +} + +var gallerySortOptions = sortOptions{ + "created_at", + "date", + "file_count", + "file_mod_time", + "id", + "images_count", + "path", + "performer_count", + "random", + "rating", + "tag_count", + "title", + "updated_at", +} + +func (qb *GalleryStore) setGallerySort(query *queryBuilder, findFilter *models.FindFilterType) error { + if findFilter == nil || findFilter.Sort == nil || *findFilter.Sort == "" { + return nil + } + sort := findFilter.GetSort("path") + direction := findFilter.GetDirection() + + // CVE-2024-32231 - ensure sort is in the list of allowed sorts + if err := gallerySortOptions.validateSort(sort); err != nil { + return err + } + + addFileTable := func() { + query.addJoins( + join{ + sort: true, + table: galleriesFilesTable, + onClause: "galleries_files.gallery_id = galleries.id AND galleries_files.\"primary\" = true", + }, + join{ + sort: true, + table: fileTable, + onClause: "galleries_files.file_id = files.id", + }, + ) + } + + addFolderTable := func() { + query.addJoins( + join{ + sort: true, + table: folderTable, + onClause: "folders.id = galleries.folder_id", + }, + join{ + sort: true, + table: folderTable, + as: "file_folder", + onClause: "files.parent_folder_id = file_folder.id", + }, + ) + } + + switch sort { + case "file_count": + query.sort += getCountSort(galleryTable, galleriesFilesTable, galleryIDColumn, direction) + case "images_count": + query.sort += getCountSort(galleryTable, galleriesImagesTable, galleryIDColumn, direction) + case "tag_count": + query.sort += getCountSort(galleryTable, galleriesTagsTable, galleryIDColumn, direction) + case "performer_count": + query.sort += getCountSort(galleryTable, performersGalleriesTable, galleryIDColumn, direction) + case "path": + // special handling for path + addFileTable() + addFolderTable() + query.sort += fmt.Sprintf(" ORDER BY COALESCE(folders.path, '') || COALESCE(file_folder.path, '') || COALESCE(files.basename, '') COLLATE NATURAL_CI %s", direction) + query.addGroupBy("folders.path", "file_folder.path", "files.basename") + case "file_mod_time": + sort = "mod_time" + addFileTable() + add, agg := getSort(sort, direction, fileTable) + query.sort += add + query.addGroupBy(agg...) + case "title": + addFileTable() + addFolderTable() + query.sort += " ORDER BY COALESCE(galleries.title, files.basename, " + getBasenameSQL("COALESCE(folders.path, '')") + ") COLLATE NATURAL_CI " + direction + ", file_folder.path COLLATE NATURAL_CI " + direction + query.addGroupBy("galleries.title", "files.basename", "folders.path", "file_folder.path") + default: + add, agg := getSort(sort, direction, "galleries") + query.sort += add + query.addGroupBy(agg...) + } + + // Whatever the sorting, always use title/id as a final sort + query.sort += ", COALESCE(galleries.title, CAST(galleries.id as text)) COLLATE NATURAL_CI ASC" + query.addGroupBy("galleries.title", "galleries.id") + + return nil +} + +func (qb *GalleryStore) GetURLs(ctx context.Context, galleryID int) ([]string, error) { + return galleriesURLsTableMgr.get(ctx, galleryID) +} + +func (qb *GalleryStore) AddFileID(ctx context.Context, id int, fileID models.FileID) error { + const firstPrimary = false + return galleriesFilesTableMgr.insertJoins(ctx, id, firstPrimary, []models.FileID{fileID}) +} + +func (qb *GalleryStore) GetPerformerIDs(ctx context.Context, id int) ([]int, error) { + return galleryRepository.performers.getIDs(ctx, id) +} + +func (qb *GalleryStore) GetTagIDs(ctx context.Context, id int) ([]int, error) { + return galleryRepository.tags.getIDs(ctx, id) +} + +func (qb *GalleryStore) GetImageIDs(ctx context.Context, galleryID int) ([]int, error) { + return galleryRepository.images.getIDs(ctx, galleryID) +} + +func (qb *GalleryStore) AddImages(ctx context.Context, galleryID int, imageIDs ...int) error { + return galleryRepository.images.insertOrIgnore(ctx, galleryID, imageIDs...) +} + +func (qb *GalleryStore) RemoveImages(ctx context.Context, galleryID int, imageIDs ...int) error { + return galleryRepository.images.destroyJoins(ctx, galleryID, imageIDs...) +} + +func (qb *GalleryStore) UpdateImages(ctx context.Context, galleryID int, imageIDs []int) error { + // Delete the existing joins and then create new ones + return galleryRepository.images.replace(ctx, galleryID, imageIDs) +} + +func (qb *GalleryStore) SetCover(ctx context.Context, galleryID int, coverImageID int) error { + return imageGalleriesTableMgr.setCover(ctx, coverImageID, galleryID) +} + +func (qb *GalleryStore) ResetCover(ctx context.Context, galleryID int) error { + return imageGalleriesTableMgr.resetCover(ctx, galleryID) +} + +func (qb *GalleryStore) GetSceneIDs(ctx context.Context, id int) ([]int, error) { + return galleryRepository.scenes.getIDs(ctx, id) +} diff --git a/pkg/postgres/gallery_chapter.go b/pkg/postgres/gallery_chapter.go new file mode 100644 index 0000000000..91599a1b40 --- /dev/null +++ b/pkg/postgres/gallery_chapter.go @@ -0,0 +1,257 @@ +package postgres + +import ( + "context" + "database/sql" + "errors" + "fmt" + "slices" + + "github.com/doug-martin/goqu/v9" + "github.com/doug-martin/goqu/v9/exp" + "github.com/jmoiron/sqlx" + + "github.com/stashapp/stash/pkg/models" +) + +const ( + galleriesChaptersTable = "galleries_chapters" +) + +type galleryChapterRow struct { + ID int `db:"id" goqu:"skipinsert"` + Title string `db:"title"` // TODO: make db schema (and gql schema) nullable + ImageIndex int `db:"image_index"` + GalleryID int `db:"gallery_id"` + CreatedAt Timestamp `db:"created_at"` + UpdatedAt Timestamp `db:"updated_at"` +} + +func (r *galleryChapterRow) fromGalleryChapter(o models.GalleryChapter) { + r.ID = o.ID + r.Title = o.Title + r.ImageIndex = o.ImageIndex + r.GalleryID = o.GalleryID + r.CreatedAt = Timestamp{Timestamp: o.CreatedAt} + r.UpdatedAt = Timestamp{Timestamp: o.UpdatedAt} +} + +func (r *galleryChapterRow) resolve() *models.GalleryChapter { + ret := &models.GalleryChapter{ + ID: r.ID, + Title: r.Title, + ImageIndex: r.ImageIndex, + GalleryID: r.GalleryID, + CreatedAt: r.CreatedAt.Timestamp.UTC(), + UpdatedAt: r.UpdatedAt.Timestamp.UTC(), + } + + return ret +} + +type galleryChapterRowRecord struct { + updateRecord +} + +func (r *galleryChapterRowRecord) fromPartial(o models.GalleryChapterPartial) { + // TODO: replace with setNullString after schema is made nullable + // r.setNullString("title", o.Title) + // saves a null input as the empty string + if o.Title.Set { + r.set("title", o.Title.Value) + } + r.setInt("image_index", o.ImageIndex) + r.setInt("gallery_id", o.GalleryID) + r.setTimestamp("created_at", o.CreatedAt) + r.setTimestamp("updated_at", o.UpdatedAt) +} + +type GalleryChapterStore struct { + repository + + tableMgr *table +} + +func NewGalleryChapterStore() *GalleryChapterStore { + return &GalleryChapterStore{ + repository: repository{ + tableName: galleriesChaptersTable, + idColumn: idColumn, + }, + tableMgr: galleriesChaptersTableMgr, + } +} + +func (qb *GalleryChapterStore) table() exp.IdentifierExpression { + return qb.tableMgr.table +} + +func (qb *GalleryChapterStore) selectDataset() *goqu.SelectDataset { + return dialect.From(qb.table()).Select(qb.table().All()) +} + +func (qb *GalleryChapterStore) Create(ctx context.Context, newObject *models.GalleryChapter) error { + var r galleryChapterRow + r.fromGalleryChapter(*newObject) + + id, err := qb.tableMgr.insertID(ctx, r) + if err != nil { + return err + } + + updated, err := qb.find(ctx, id) + if err != nil { + return fmt.Errorf("finding after create: %w", err) + } + + *newObject = *updated + + return nil +} + +func (qb *GalleryChapterStore) Update(ctx context.Context, updatedObject *models.GalleryChapter) error { + var r galleryChapterRow + r.fromGalleryChapter(*updatedObject) + + if err := qb.tableMgr.updateByID(ctx, updatedObject.ID, r); err != nil { + return err + } + + return nil +} + +func (qb *GalleryChapterStore) UpdatePartial(ctx context.Context, id int, partial models.GalleryChapterPartial) (*models.GalleryChapter, error) { + r := galleryChapterRowRecord{ + updateRecord{ + Record: make(exp.Record), + }, + } + + r.fromPartial(partial) + + if len(r.Record) > 0 { + if err := qb.tableMgr.updateByID(ctx, id, r.Record); err != nil { + return nil, err + } + } + + return qb.find(ctx, id) +} + +func (qb *GalleryChapterStore) Destroy(ctx context.Context, id int) error { + return qb.destroyExisting(ctx, []int{id}) +} + +// returns nil, nil if not found +func (qb *GalleryChapterStore) Find(ctx context.Context, id int) (*models.GalleryChapter, error) { + ret, err := qb.find(ctx, id) + if errors.Is(err, sql.ErrNoRows) { + return nil, nil + } + return ret, err +} + +func (qb *GalleryChapterStore) FindMany(ctx context.Context, ids []int) ([]*models.GalleryChapter, error) { + ret := make([]*models.GalleryChapter, len(ids)) + + if len(ids) == 0 { + return ret, nil + } + + table := qb.table() + q := qb.selectDataset().Prepared(true).Where(table.Col(idColumn).In(ids)) + unsorted, err := qb.getMany(ctx, q) + if err != nil { + return nil, err + } + + for _, s := range unsorted { + i := slices.Index(ids, s.ID) + ret[i] = s + } + + for i := range ret { + if ret[i] == nil { + return nil, fmt.Errorf("gallery chapter with id %d not found", ids[i]) + } + } + + return ret, nil +} + +// returns nil, sql.ErrNoRows if not found +func (qb *GalleryChapterStore) find(ctx context.Context, id int) (*models.GalleryChapter, error) { + q := qb.selectDataset().Where(qb.tableMgr.byID(id)) + + ret, err := qb.get(ctx, q) + if err != nil { + return nil, err + } + + return ret, nil +} + +// returns nil, sql.ErrNoRows if not found +func (qb *GalleryChapterStore) get(ctx context.Context, q *goqu.SelectDataset) (*models.GalleryChapter, error) { + ret, err := qb.getMany(ctx, q) + if err != nil { + return nil, err + } + + if len(ret) == 0 { + return nil, sql.ErrNoRows + } + + return ret[0], nil +} + +func (qb *GalleryChapterStore) getMany(ctx context.Context, q *goqu.SelectDataset) ([]*models.GalleryChapter, error) { + const single = false + var ret []*models.GalleryChapter + if err := queryFunc(ctx, q, single, func(r *sqlx.Rows) error { + var f galleryChapterRow + if err := r.StructScan(&f); err != nil { + return err + } + + s := f.resolve() + + ret = append(ret, s) + return nil + }); err != nil { + return nil, err + } + + return ret, nil +} + +func (qb *GalleryChapterStore) FindByGalleryID(ctx context.Context, galleryID int) ([]*models.GalleryChapter, error) { + query := ` + SELECT galleries_chapters.* FROM galleries_chapters + WHERE galleries_chapters.gallery_id = ? + GROUP BY galleries_chapters.id + ORDER BY galleries_chapters.image_index ASC + ` + args := []interface{}{galleryID} + return qb.queryGalleryChapters(ctx, query, args) +} + +func (qb *GalleryChapterStore) queryGalleryChapters(ctx context.Context, query string, args []interface{}) ([]*models.GalleryChapter, error) { + const single = false + var ret []*models.GalleryChapter + if err := qb.queryFunc(ctx, query, args, single, func(r *sqlx.Rows) error { + var f galleryChapterRow + if err := r.StructScan(&f); err != nil { + return err + } + + s := f.resolve() + + ret = append(ret, s) + return nil + }); err != nil { + return nil, err + } + + return ret, nil +} diff --git a/pkg/postgres/gallery_filter.go b/pkg/postgres/gallery_filter.go new file mode 100644 index 0000000000..d1737128d0 --- /dev/null +++ b/pkg/postgres/gallery_filter.go @@ -0,0 +1,464 @@ +package postgres + +import ( + "context" + "fmt" + "path/filepath" + "regexp" + + "github.com/stashapp/stash/pkg/models" +) + +type galleryFilterHandler struct { + galleryFilter *models.GalleryFilterType +} + +func (qb *galleryFilterHandler) validate() error { + galleryFilter := qb.galleryFilter + if galleryFilter == nil { + return nil + } + + if err := validateFilterCombination(galleryFilter.OperatorFilter); err != nil { + return err + } + + if subFilter := galleryFilter.SubFilter(); subFilter != nil { + sqb := &galleryFilterHandler{galleryFilter: subFilter} + if err := sqb.validate(); err != nil { + return err + } + } + + return nil +} + +func (qb *galleryFilterHandler) handle(ctx context.Context, f *filterBuilder) { + galleryFilter := qb.galleryFilter + if galleryFilter == nil { + return + } + + if err := qb.validate(); err != nil { + f.setError(err) + return + } + + sf := galleryFilter.SubFilter() + if sf != nil { + sub := &galleryFilterHandler{sf} + handleSubFilter(ctx, sub, f, galleryFilter.OperatorFilter) + } + + f.handleCriterion(ctx, qb.criterionHandler()) +} + +func (qb *galleryFilterHandler) criterionHandler() criterionHandler { + filter := qb.galleryFilter + return compoundHandler{ + intCriterionHandler(filter.ID, "galleries.id", nil), + stringCriterionHandler(filter.Title, "galleries.title"), + stringCriterionHandler(filter.Code, "galleries.code"), + stringCriterionHandler(filter.Details, "galleries.details"), + stringCriterionHandler(filter.Photographer, "galleries.photographer"), + + criterionHandlerFunc(func(ctx context.Context, f *filterBuilder) { + if filter.Checksum != nil { + galleryRepository.addGalleriesFilesTable(f) + f.addLeftJoin(fingerprintTable, "fingerprints_md5", "galleries_files.file_id = fingerprints_md5.file_id AND fingerprints_md5.type = 'md5'") + } + + stringCriterionHandler(filter.Checksum, "fingerprints_md5.fingerprint")(ctx, f) + }), + + criterionHandlerFunc(func(ctx context.Context, f *filterBuilder) { + if filter.IsZip != nil { + galleryRepository.addGalleriesFilesTable(f) + if *filter.IsZip { + + f.addWhere("galleries_files.file_id IS NOT NULL") + } else { + f.addWhere("galleries_files.file_id IS NULL") + } + } + }), + + qb.pathCriterionHandler(filter.Path), + qb.fileCountCriterionHandler(filter.FileCount), + intCriterionHandler(filter.Rating100, "galleries.rating", nil), + qb.urlsCriterionHandler(filter.URL), + boolCriterionHandler(filter.Organized, "galleries.organized", nil), + qb.missingCriterionHandler(filter.IsMissing), + qb.tagsCriterionHandler(filter.Tags), + qb.tagCountCriterionHandler(filter.TagCount), + qb.performersCriterionHandler(filter.Performers), + qb.performerCountCriterionHandler(filter.PerformerCount), + qb.scenesCriterionHandler(filter.Scenes), + qb.hasChaptersCriterionHandler(filter.HasChapters), + studioCriterionHandler(galleryTable, filter.Studios), + qb.performerTagsCriterionHandler(filter.PerformerTags), + qb.averageResolutionCriterionHandler(filter.AverageResolution), + qb.imageCountCriterionHandler(filter.ImageCount), + qb.performerFavoriteCriterionHandler(filter.PerformerFavorite), + qb.performerAgeCriterionHandler(filter.PerformerAge), + &dateCriterionHandler{filter.Date, "galleries.date", nil}, + ×tampCriterionHandler{filter.CreatedAt, "galleries.created_at", nil}, + ×tampCriterionHandler{filter.UpdatedAt, "galleries.updated_at", nil}, + + &relatedFilterHandler{ + relatedIDCol: "scenes_galleries.scene_id", + relatedRepo: sceneRepository.repository, + relatedHandler: &sceneFilterHandler{filter.ScenesFilter}, + joinFn: func(f *filterBuilder) { + galleryRepository.scenes.innerJoin(f, "", "galleries.id") + }, + }, + + &relatedFilterHandler{ + relatedIDCol: "galleries_images.image_id", + relatedRepo: imageRepository.repository, + relatedHandler: &imageFilterHandler{filter.ImagesFilter}, + joinFn: func(f *filterBuilder) { + galleryRepository.images.innerJoin(f, "", "galleries.id") + }, + }, + + &relatedFilterHandler{ + relatedIDCol: "performers_join.performer_id", + relatedRepo: performerRepository.repository, + relatedHandler: &performerFilterHandler{filter.PerformersFilter}, + joinFn: func(f *filterBuilder) { + galleryRepository.performers.innerJoin(f, "performers_join", "galleries.id") + }, + }, + + &relatedFilterHandler{ + relatedIDCol: "galleries.studio_id", + relatedRepo: studioRepository.repository, + relatedHandler: &studioFilterHandler{filter.StudiosFilter}, + }, + + &relatedFilterHandler{ + relatedIDCol: "gallery_tag.tag_id", + relatedRepo: tagRepository.repository, + relatedHandler: &tagFilterHandler{filter.TagsFilter}, + joinFn: func(f *filterBuilder) { + galleryRepository.tags.innerJoin(f, "gallery_tag", "galleries.id") + }, + }, + + &relatedFilterHandler{ + relatedIDCol: "files.id", + relatedRepo: fileRepository.repository, + relatedHandler: &fileFilterHandler{ + fileFilter: filter.FilesFilter, + isRelated: true, + }, + joinFn: func(f *filterBuilder) { + galleryRepository.addFilesTable(f) + galleryRepository.addFoldersTable(f) + }, + // don't use a subquery; join directly + directJoin: true, + }, + + &relatedFilterHandler{ + relatedIDCol: "gallery_folder.id", + relatedRepo: folderRepository.repository, + relatedHandler: &folderFilterHandler{ + folderFilter: filter.FoldersFilter, + table: "gallery_folder", + isRelated: true, + }, + joinFn: func(f *filterBuilder) { + f.addLeftJoin(folderTable, "gallery_folder", "galleries.folder_id = gallery_folder.id") + }, + // don't use a subquery; join directly + directJoin: true, + }, + } +} + +func (qb *galleryFilterHandler) urlsCriterionHandler(url *models.StringCriterionInput) criterionHandlerFunc { + h := stringListCriterionHandlerBuilder{ + primaryTable: galleryTable, + primaryFK: galleryIDColumn, + joinTable: galleriesURLsTable, + stringColumn: galleriesURLColumn, + addJoinTable: func(f *filterBuilder) { + galleriesURLsTableMgr.join(f, "", "galleries.id") + }, + } + + return h.handler(url) +} + +func (qb *galleryFilterHandler) getMultiCriterionHandlerBuilder(foreignTable, joinTable, foreignFK string, addJoinsFunc func(f *filterBuilder)) multiCriterionHandlerBuilder { + return multiCriterionHandlerBuilder{ + primaryTable: galleryTable, + foreignTable: foreignTable, + joinTable: joinTable, + primaryFK: galleryIDColumn, + foreignFK: foreignFK, + addJoinsFunc: addJoinsFunc, + } +} + +func (qb *galleryFilterHandler) pathCriterionHandler(c *models.StringCriterionInput) criterionHandlerFunc { + return func(ctx context.Context, f *filterBuilder) { + if c != nil { + galleryRepository.addFoldersTable(f) + f.addLeftJoin(folderTable, "gallery_folder", "galleries.folder_id = gallery_folder.id") + + const pathColumn = "folders.path" + const basenameColumn = "files.basename" + const folderPathColumn = "gallery_folder.path" + + addWildcards := true + not := false + + if modifier := c.Modifier; c.Modifier.IsValid() { + switch modifier { + case models.CriterionModifierIncludes: + clause := getPathSearchClauseMany(pathColumn, basenameColumn, c.Value, addWildcards, not) + clause2 := getStringSearchClause([]string{folderPathColumn}, c.Value, false) + f.whereClauses = append(f.whereClauses, orClauses(clause, clause2)) + case models.CriterionModifierExcludes: + not = true + clause := getPathSearchClauseMany(pathColumn, basenameColumn, c.Value, addWildcards, not) + clause2 := getStringSearchClause([]string{folderPathColumn}, c.Value, true) + f.whereClauses = append(f.whereClauses, orClauses(clause, clause2)) + case models.CriterionModifierEquals: + addWildcards = false + clause := getPathSearchClause(pathColumn, basenameColumn, c.Value, addWildcards, not) + clause2 := makeClause(folderPathColumn+" ILIKE ?", c.Value) + f.whereClauses = append(f.whereClauses, orClauses(clause, clause2)) + case models.CriterionModifierNotEquals: + addWildcards = false + not = true + clause := getPathSearchClause(pathColumn, basenameColumn, c.Value, addWildcards, not) + clause2 := makeClause(folderPathColumn+" NOT ILIKE ?", c.Value) + f.whereClauses = append(f.whereClauses, orClauses(clause, clause2)) + case models.CriterionModifierMatchesRegex: + if _, err := regexp.Compile(c.Value); err != nil { + f.setError(err) + return + } + filepathColumn := fmt.Sprintf("%s || '%s' || %s", pathColumn, string(filepath.Separator), basenameColumn) + clause := makeClause(fmt.Sprintf("%s IS NOT NULL AND %s IS NOT NULL AND regex_match(%s, ?)", pathColumn, basenameColumn, filepathColumn), c.Value) + clause2 := makeClause(fmt.Sprintf("%s IS NOT NULL AND regex_match(%[1]s, ?)", folderPathColumn), c.Value) + f.whereClauses = append(f.whereClauses, orClauses(clause, clause2)) + case models.CriterionModifierNotMatchesRegex: + if _, err := regexp.Compile(c.Value); err != nil { + f.setError(err) + return + } + filepathColumn := fmt.Sprintf("%s || '%s' || %s", pathColumn, string(filepath.Separator), basenameColumn) + f.addWhere(fmt.Sprintf("%s IS NULL OR %s IS NULL OR NOT regex_match(%s, ?)", pathColumn, basenameColumn, filepathColumn), c.Value) + f.addWhere(fmt.Sprintf("%s IS NULL OR NOT regex_match(%[1]s, ?)", folderPathColumn), c.Value) + case models.CriterionModifierIsNull: + f.addWhere(fmt.Sprintf("%s IS NULL OR TRIM(%[1]s) = '' OR %s IS NULL OR TRIM(%[2]s) = ''", pathColumn, basenameColumn)) + f.addWhere(fmt.Sprintf("%s IS NULL OR TRIM(%[1]s) = ''", folderPathColumn)) + case models.CriterionModifierNotNull: + clause := makeClause(fmt.Sprintf("%s IS NOT NULL AND TRIM(%[1]s) != '' AND %s IS NOT NULL AND TRIM(%[2]s) != ''", pathColumn, basenameColumn)) + clause2 := makeClause(fmt.Sprintf("%s IS NOT NULL AND TRIM(%[1]s) != ''", folderPathColumn)) + f.whereClauses = append(f.whereClauses, orClauses(clause, clause2)) + default: + panic("unsupported string filter modifier") + } + } + } + } +} + +func (qb *galleryFilterHandler) fileCountCriterionHandler(fileCount *models.IntCriterionInput) criterionHandlerFunc { + h := countCriterionHandlerBuilder{ + primaryTable: galleryTable, + joinTable: galleriesFilesTable, + primaryFK: galleryIDColumn, + } + + return h.handler(fileCount) +} + +func (qb *galleryFilterHandler) missingCriterionHandler(isMissing *string) criterionHandlerFunc { + return func(ctx context.Context, f *filterBuilder) { + if isMissing != nil && *isMissing != "" { + switch *isMissing { + case "url": + galleriesURLsTableMgr.join(f, "", "galleries.id") + f.addWhere("gallery_urls.url IS NULL") + case "scenes": + f.addLeftJoin("scenes_galleries", "scenes_join", "scenes_join.gallery_id = galleries.id") + f.addWhere("scenes_join.gallery_id IS NULL") + case "studio": + f.addWhere("galleries.studio_id IS NULL") + case "performers": + galleryRepository.performers.join(f, "performers_join", "galleries.id") + f.addWhere("performers_join.gallery_id IS NULL") + case "date": + f.addWhere("galleries.date IS NULL") + case "tags": + galleryRepository.tags.join(f, "tags_join", "galleries.id") + f.addWhere("tags_join.gallery_id IS NULL") + default: + f.addWhere("(galleries." + *isMissing + " IS NULL OR TRIM(CAST(galleries." + *isMissing + " AS TEXT)) = '')") + } + } + } +} + +func (qb *galleryFilterHandler) tagsCriterionHandler(tags *models.HierarchicalMultiCriterionInput) criterionHandlerFunc { + h := joinedHierarchicalMultiCriterionHandlerBuilder{ + primaryTable: galleryTable, + foreignTable: tagTable, + foreignFK: "tag_id", + + relationsTable: "tags_relations", + joinAs: "gallery_tag", + joinTable: galleriesTagsTable, + primaryFK: galleryIDColumn, + } + + return h.handler(tags) +} + +func (qb *galleryFilterHandler) tagCountCriterionHandler(tagCount *models.IntCriterionInput) criterionHandlerFunc { + h := countCriterionHandlerBuilder{ + primaryTable: galleryTable, + joinTable: galleriesTagsTable, + primaryFK: galleryIDColumn, + } + + return h.handler(tagCount) +} + +func (qb *galleryFilterHandler) scenesCriterionHandler(scenes *models.MultiCriterionInput) criterionHandlerFunc { + addJoinsFunc := func(f *filterBuilder) { + galleryRepository.scenes.join(f, "", "galleries.id") + f.addLeftJoin("scenes", "", "scenes_galleries.scene_id = scenes.id") + } + h := qb.getMultiCriterionHandlerBuilder(sceneTable, galleriesScenesTable, "scene_id", addJoinsFunc) + return h.handler(scenes) +} + +func (qb *galleryFilterHandler) performersCriterionHandler(performers *models.MultiCriterionInput) criterionHandlerFunc { + h := joinedMultiCriterionHandlerBuilder{ + primaryTable: galleryTable, + joinTable: performersGalleriesTable, + joinAs: "performers_join", + primaryFK: galleryIDColumn, + foreignFK: performerIDColumn, + + addJoinTable: func(f *filterBuilder) { + galleryRepository.performers.join(f, "performers_join", "galleries.id") + }, + } + + return h.handler(performers) +} + +func (qb *galleryFilterHandler) performerCountCriterionHandler(performerCount *models.IntCriterionInput) criterionHandlerFunc { + h := countCriterionHandlerBuilder{ + primaryTable: galleryTable, + joinTable: performersGalleriesTable, + primaryFK: galleryIDColumn, + } + + return h.handler(performerCount) +} + +func (qb *galleryFilterHandler) imageCountCriterionHandler(imageCount *models.IntCriterionInput) criterionHandlerFunc { + h := countCriterionHandlerBuilder{ + primaryTable: galleryTable, + joinTable: galleriesImagesTable, + primaryFK: galleryIDColumn, + } + + return h.handler(imageCount) +} + +func (qb *galleryFilterHandler) hasChaptersCriterionHandler(hasChapters *string) criterionHandlerFunc { + return func(ctx context.Context, f *filterBuilder) { + if hasChapters != nil { + f.addLeftJoin("galleries_chapters", "", "galleries_chapters.gallery_id = galleries.id") + if *hasChapters == "true" { + f.addHaving("count(galleries_chapters.gallery_id) > 0") + } else { + f.addWhere("galleries_chapters.id IS NULL") + } + } + } +} + +func (qb *galleryFilterHandler) performerTagsCriterionHandler(tags *models.HierarchicalMultiCriterionInput) criterionHandler { + return &joinedPerformerTagsHandler{ + criterion: tags, + primaryTable: galleryTable, + joinTable: performersGalleriesTable, + joinPrimaryKey: galleryIDColumn, + } +} + +func (qb *galleryFilterHandler) performerFavoriteCriterionHandler(performerfavorite *bool) criterionHandlerFunc { + return func(ctx context.Context, f *filterBuilder) { + if performerfavorite != nil { + f.addLeftJoin("performers_galleries", "", "galleries.id = performers_galleries.gallery_id") + + if *performerfavorite { + // contains at least one favorite + f.addLeftJoin("performers", "", "performers.id = performers_galleries.performer_id") + f.addWhere("performers.favorite = true") + } else { + // contains zero favorites + f.addLeftJoin(`(SELECT performers_galleries.gallery_id as id FROM performers_galleries +JOIN performers ON performers.id = performers_galleries.performer_id +GROUP BY performers_galleries.gallery_id HAVING SUM(performers.favorite) = false)`, "nofaves", "galleries.id = nofaves.id") + f.addWhere("performers_galleries.gallery_id IS NULL OR nofaves.id IS NOT NULL") + } + } + } +} + +func (qb *galleryFilterHandler) performerAgeCriterionHandler(performerAge *models.IntCriterionInput) criterionHandlerFunc { + return func(ctx context.Context, f *filterBuilder) { + if performerAge != nil { + f.addInnerJoin("performers_galleries", "", "galleries.id = performers_galleries.gallery_id") + f.addInnerJoin("performers", "", "performers_galleries.performer_id = performers.id") + + f.addWhere("galleries.date != '' AND performers.birthdate != ''") + f.addWhere("galleries.date IS NOT NULL AND performers.birthdate IS NOT NULL") + + ageCalc := "EXTRACT(YEAR FROM AGE(galleries.date, performers.birthdate))" + whereClause, args := getIntWhereClause(ageCalc, performerAge.Modifier, performerAge.Value, performerAge.Value2) + f.addWhere(whereClause, args...) + } + } +} + +func (qb *galleryFilterHandler) averageResolutionCriterionHandler(resolution *models.ResolutionCriterionInput) criterionHandlerFunc { + return func(ctx context.Context, f *filterBuilder) { + if resolution != nil && resolution.Value.IsValid() { + galleryRepository.images.join(f, "images_join", "galleries.id") + f.addLeftJoin("images", "", "images_join.image_id = images.id") + f.addLeftJoin("images_files", "", "images.id = images_files.image_id") + f.addLeftJoin("image_files", "", "images_files.file_id = image_files.file_id") + + mn := resolution.Value.GetMinResolution() + mx := resolution.Value.GetMaxResolution() + + const widthHeight = "avg(LEAST(image_files.width, image_files.height))" + + switch resolution.Modifier { + case models.CriterionModifierEquals: + f.addHaving(fmt.Sprintf("%s BETWEEN %d AND %d", widthHeight, mn, mx)) + case models.CriterionModifierNotEquals: + f.addHaving(fmt.Sprintf("%s NOT BETWEEN %d AND %d", widthHeight, mn, mx)) + case models.CriterionModifierLessThan: + f.addHaving(fmt.Sprintf("%s < %d", widthHeight, mn)) + case models.CriterionModifierGreaterThan: + f.addHaving(fmt.Sprintf("%s > %d", widthHeight, mx)) + } + } + } +} diff --git a/pkg/postgres/group.go b/pkg/postgres/group.go new file mode 100644 index 0000000000..b2a94f4e31 --- /dev/null +++ b/pkg/postgres/group.go @@ -0,0 +1,724 @@ +package postgres + +import ( + "context" + "database/sql" + "errors" + "fmt" + "slices" + + "github.com/doug-martin/goqu/v9" + "github.com/doug-martin/goqu/v9/exp" + "github.com/jmoiron/sqlx" + "gopkg.in/guregu/null.v4" + "gopkg.in/guregu/null.v4/zero" + + "github.com/stashapp/stash/pkg/models" +) + +const ( + groupTable = "groups" + groupIDColumn = "group_id" + + groupFrontImageBlobColumn = "front_image_blob" + groupBackImageBlobColumn = "back_image_blob" + + groupsTagsTable = "groups_tags" + + groupURLsTable = "group_urls" + groupURLColumn = "url" + + groupRelationsTable = "groups_relations" +) + +type groupRow struct { + ID int `db:"id" goqu:"skipinsert"` + Name zero.String `db:"name"` + Aliases zero.String `db:"aliases"` + Duration null.Int `db:"duration"` + Date NullDate `db:"date"` + DatePrecision null.Int `db:"date_precision"` + // expressed as 1-100 + Rating null.Int `db:"rating"` + StudioID null.Int `db:"studio_id,omitempty"` + Director zero.String `db:"director"` + Description zero.String `db:"description"` + CreatedAt Timestamp `db:"created_at"` + UpdatedAt Timestamp `db:"updated_at"` + + // not used in resolutions or updates + FrontImageBlob zero.String `db:"front_image_blob"` + BackImageBlob zero.String `db:"back_image_blob"` +} + +func (r *groupRow) fromGroup(o models.Group) { + r.ID = o.ID + r.Name = zero.StringFrom(o.Name) + r.Aliases = zero.StringFrom(o.Aliases) + r.Duration = intFromPtr(o.Duration) + r.Date = NullDateFromDatePtr(o.Date) + r.DatePrecision = datePrecisionFromDatePtr(o.Date) + r.Rating = intFromPtr(o.Rating) + r.StudioID = intFromPtr(o.StudioID) + r.Director = zero.StringFrom(o.Director) + r.Description = zero.StringFrom(o.Synopsis) + r.CreatedAt = Timestamp{Timestamp: o.CreatedAt} + r.UpdatedAt = Timestamp{Timestamp: o.UpdatedAt} +} + +func (r *groupRow) resolve() *models.Group { + ret := &models.Group{ + ID: r.ID, + Name: r.Name.String, + Aliases: r.Aliases.String, + Duration: nullIntPtr(r.Duration), + Date: r.Date.DatePtr(r.DatePrecision), + Rating: nullIntPtr(r.Rating), + StudioID: nullIntPtr(r.StudioID), + Director: r.Director.String, + Synopsis: r.Description.String, + CreatedAt: r.CreatedAt.Timestamp.UTC(), + UpdatedAt: r.UpdatedAt.Timestamp.UTC(), + } + + return ret +} + +type groupRowRecord struct { + updateRecord +} + +func (r *groupRowRecord) fromPartial(o models.GroupPartial) { + r.setNullString("name", o.Name) + r.setNullString("aliases", o.Aliases) + r.setNullInt("duration", o.Duration) + r.setNullDate("date", "date_precision", o.Date) + r.setNullInt("rating", o.Rating) + r.setNullInt("studio_id", o.StudioID) + r.setNullString("director", o.Director) + r.setNullString("description", o.Synopsis) + r.setTimestamp("created_at", o.CreatedAt) + r.setTimestamp("updated_at", o.UpdatedAt) +} + +type groupRepositoryType struct { + repository + scenes repository + tags joinRepository +} + +var ( + groupRepository = groupRepositoryType{ + repository: repository{ + tableName: groupTable, + idColumn: idColumn, + }, + scenes: repository{ + tableName: groupsScenesTable, + idColumn: groupIDColumn, + }, + tags: joinRepository{ + repository: repository{ + tableName: groupsTagsTable, + idColumn: groupIDColumn, + }, + fkColumn: tagIDColumn, + foreignTable: tagTable, + orderBy: tagTableSortSQL, + }, + } +) + +type GroupStore struct { + blobJoinQueryBuilder + tagRelationshipStore + groupRelationshipStore + + tableMgr *table +} + +func NewGroupStore(blobStore *BlobStore) *GroupStore { + return &GroupStore{ + blobJoinQueryBuilder: blobJoinQueryBuilder{ + blobStore: blobStore, + joinTable: groupTable, + }, + tagRelationshipStore: tagRelationshipStore{ + idRelationshipStore: idRelationshipStore{ + joinTable: groupsTagsTableMgr, + }, + }, + groupRelationshipStore: groupRelationshipStore{ + table: groupRelationshipTableMgr, + }, + + tableMgr: groupTableMgr, + } +} + +func (qb *GroupStore) table() exp.IdentifierExpression { + return qb.tableMgr.table +} + +func (qb *GroupStore) selectDataset() *goqu.SelectDataset { + return dialect.From(qb.table()).Select(qb.table().All()) +} + +func (qb *GroupStore) Create(ctx context.Context, newObject *models.Group) error { + var r groupRow + r.fromGroup(*newObject) + + id, err := qb.tableMgr.insertID(ctx, r) + if err != nil { + return err + } + + if newObject.URLs.Loaded() { + const startPos = 0 + if err := groupsURLsTableMgr.insertJoins(ctx, id, startPos, newObject.URLs.List()); err != nil { + return err + } + } + + if err := qb.tagRelationshipStore.createRelationships(ctx, id, newObject.TagIDs); err != nil { + return err + } + + if err := qb.groupRelationshipStore.createContainingRelationships(ctx, id, newObject.ContainingGroups); err != nil { + return err + } + + if err := qb.groupRelationshipStore.createSubRelationships(ctx, id, newObject.SubGroups); err != nil { + return err + } + + updated, err := qb.find(ctx, id) + if err != nil { + return fmt.Errorf("finding after create: %w", err) + } + + *newObject = *updated + + return nil +} + +func (qb *GroupStore) UpdatePartial(ctx context.Context, id int, partial models.GroupPartial) (*models.Group, error) { + r := groupRowRecord{ + updateRecord{ + Record: make(exp.Record), + }, + } + + r.fromPartial(partial) + + if len(r.Record) > 0 { + if err := qb.tableMgr.updateByID(ctx, id, r.Record); err != nil { + return nil, err + } + } + + if partial.URLs != nil { + if err := groupsURLsTableMgr.modifyJoins(ctx, id, partial.URLs.Values, partial.URLs.Mode); err != nil { + return nil, err + } + } + + if err := qb.tagRelationshipStore.modifyRelationships(ctx, id, partial.TagIDs); err != nil { + return nil, err + } + + if err := qb.groupRelationshipStore.modifyContainingRelationships(ctx, id, partial.ContainingGroups); err != nil { + return nil, err + } + + if err := qb.groupRelationshipStore.modifySubRelationships(ctx, id, partial.SubGroups); err != nil { + return nil, err + } + + return qb.find(ctx, id) +} + +func (qb *GroupStore) Update(ctx context.Context, updatedObject *models.Group) error { + var r groupRow + r.fromGroup(*updatedObject) + + if err := qb.tableMgr.updateByID(ctx, updatedObject.ID, r); err != nil { + return err + } + + if updatedObject.URLs.Loaded() { + if err := groupsURLsTableMgr.replaceJoins(ctx, updatedObject.ID, updatedObject.URLs.List()); err != nil { + return err + } + } + + if err := qb.tagRelationshipStore.replaceRelationships(ctx, updatedObject.ID, updatedObject.TagIDs); err != nil { + return err + } + + if err := qb.groupRelationshipStore.replaceContainingRelationships(ctx, updatedObject.ID, updatedObject.ContainingGroups); err != nil { + return err + } + + if err := qb.groupRelationshipStore.replaceSubRelationships(ctx, updatedObject.ID, updatedObject.SubGroups); err != nil { + return err + } + + return nil +} + +func (qb *GroupStore) Destroy(ctx context.Context, id int) error { + // must handle image checksums manually + if err := qb.destroyImages(ctx, id); err != nil { + return err + } + + return groupRepository.destroyExisting(ctx, []int{id}) +} + +// returns nil, nil if not found +func (qb *GroupStore) Find(ctx context.Context, id int) (*models.Group, error) { + ret, err := qb.find(ctx, id) + if errors.Is(err, sql.ErrNoRows) { + return nil, nil + } + return ret, err +} + +func (qb *GroupStore) FindMany(ctx context.Context, ids []int) ([]*models.Group, error) { + ret := make([]*models.Group, len(ids)) + + if len(ids) == 0 { + return ret, nil + } + + table := qb.table() + if err := batchExec(ids, defaultBatchSize, func(batch []int) error { + q := qb.selectDataset().Prepared(true).Where(table.Col(idColumn).In(batch)) + unsorted, err := qb.getMany(ctx, q) + if err != nil { + return err + } + + for _, s := range unsorted { + i := slices.Index(ids, s.ID) + ret[i] = s + } + + return nil + }); err != nil { + return nil, err + } + + for i := range ret { + if ret[i] == nil { + return nil, fmt.Errorf("group with id %d not found", ids[i]) + } + } + + return ret, nil +} + +// returns nil, sql.ErrNoRows if not found +func (qb *GroupStore) find(ctx context.Context, id int) (*models.Group, error) { + q := qb.selectDataset().Where(qb.tableMgr.byID(id)) + + ret, err := qb.get(ctx, q) + if err != nil { + return nil, err + } + + return ret, nil +} + +// returns nil, sql.ErrNoRows if not found +func (qb *GroupStore) get(ctx context.Context, q *goqu.SelectDataset) (*models.Group, error) { + ret, err := qb.getMany(ctx, q) + if err != nil { + return nil, err + } + + if len(ret) == 0 { + return nil, sql.ErrNoRows + } + + return ret[0], nil +} + +func (qb *GroupStore) getMany(ctx context.Context, q *goqu.SelectDataset) ([]*models.Group, error) { + const single = false + var ret []*models.Group + if err := queryFunc(ctx, q, single, func(r *sqlx.Rows) error { + var f groupRow + if err := r.StructScan(&f); err != nil { + return err + } + + s := f.resolve() + + ret = append(ret, s) + return nil + }); err != nil { + return nil, err + } + + return ret, nil +} + +func (qb *GroupStore) FindByName(ctx context.Context, name string, nocase bool) (*models.Group, error) { + // query := "SELECT * FROM groups WHERE name = ?" + // if nocase { + // query += " COLLATE NOCASE" + // } + // query += " LIMIT 1" + where := "name = ?" + if nocase { + where += " COLLATE NOCASE" + } + sq := qb.selectDataset().Prepared(true).Where(goqu.L(where, name)).Limit(1) + ret, err := qb.get(ctx, sq) + + if err != nil && !errors.Is(err, sql.ErrNoRows) { + return nil, err + } + + return ret, nil +} + +func (qb *GroupStore) FindByNames(ctx context.Context, names []string, nocase bool) ([]*models.Group, error) { + // query := "SELECT * FROM groups WHERE name" + // if nocase { + // query += " COLLATE NOCASE" + // } + // query += " IN " + getInBinding(len(names)) + where := "name" + if nocase { + where += " COLLATE NOCASE" + } + where += " IN " + getInBinding(len(names)) + var args []interface{} + for _, name := range names { + args = append(args, name) + } + sq := qb.selectDataset().Prepared(true).Where(goqu.L(where, args...)) + ret, err := qb.getMany(ctx, sq) + + if err != nil { + return nil, err + } + + return ret, nil +} + +func (qb *GroupStore) Count(ctx context.Context) (int, error) { + q := dialect.Select(goqu.COUNT("*")).From(qb.table()) + return count(ctx, q) +} + +func (qb *GroupStore) All(ctx context.Context) ([]*models.Group, error) { + table := qb.table() + + return qb.getMany(ctx, qb.selectDataset().Order( + table.Col("name").Asc(), + table.Col(idColumn).Asc(), + )) +} + +func (qb *GroupStore) makeQuery(ctx context.Context, groupFilter *models.GroupFilterType, findFilter *models.FindFilterType) (*queryBuilder, error) { + if findFilter == nil { + findFilter = &models.FindFilterType{} + } + if groupFilter == nil { + groupFilter = &models.GroupFilterType{} + } + + query := groupRepository.newQuery() + distinctIDs(&query, groupTable) + + if q := findFilter.Q; q != nil && *q != "" { + searchColumns := []string{"groups.name", "groups.aliases"} + query.parseQueryString(searchColumns, *q) + } + + filter := filterBuilderFromHandler(ctx, &groupFilterHandler{ + groupFilter: groupFilter, + }) + + if err := query.addFilter(filter); err != nil { + return nil, err + } + + if err := qb.setGroupSort(&query, findFilter); err != nil { + return nil, err + } + + query.pagination = getPagination(findFilter) + + return &query, nil +} + +func (qb *GroupStore) Query(ctx context.Context, groupFilter *models.GroupFilterType, findFilter *models.FindFilterType) ([]*models.Group, int, error) { + query, err := qb.makeQuery(ctx, groupFilter, findFilter) + if err != nil { + return nil, 0, err + } + + idsResult, countResult, err := query.executeFind(ctx) + if err != nil { + return nil, 0, err + } + + groups, err := qb.FindMany(ctx, idsResult) + if err != nil { + return nil, 0, err + } + + return groups, countResult, nil +} + +func (qb *GroupStore) QueryCount(ctx context.Context, groupFilter *models.GroupFilterType, findFilter *models.FindFilterType) (int, error) { + query, err := qb.makeQuery(ctx, groupFilter, findFilter) + if err != nil { + return 0, err + } + + return query.executeCount(ctx) +} + +var groupSortOptions = sortOptions{ + "created_at", + "date", + "duration", + "id", + "name", + "random", + "rating", + "scenes_count", + "o_counter", + "sub_group_order", + "tag_count", + "updated_at", +} + +func (qb *GroupStore) setGroupSort(query *queryBuilder, findFilter *models.FindFilterType) error { + var sort string + var direction string + if findFilter == nil { + sort = "name" + direction = "ASC" + } else { + sort = findFilter.GetSort("name") + direction = findFilter.GetDirection() + } + + // CVE-2024-32231 - ensure sort is in the list of allowed sorts + if err := groupSortOptions.validateSort(sort); err != nil { + return err + } + + switch sort { + case "sub_group_order": + // sub_group_order is a special sort that sorts by the order_index of the subgroups + if query.hasJoin("groups_parents") { + add, agg := getSort("order_index", direction, "groups_parents") + query.sort += add + query.addGroupBy(agg...) + } else { + // this will give unexpected results if the query is not filtered by a parent group and + // the group has multiple parents and order indexes + query.joinSort(groupRelationsTable, "", "groups.id = groups_relations.sub_id") + add, agg := getSort("order_index", direction, groupRelationsTable) + query.sort += add + query.addGroupBy(agg...) + } + case "tag_count": + query.sort += getCountSort(groupTable, groupsTagsTable, groupIDColumn, direction) + case "scenes_count": // generic getSort won't work for this + query.sort += getCountSort(groupTable, groupsScenesTable, groupIDColumn, direction) + case "o_counter": + query.sort += qb.sortByOCounter(direction) + default: + add, agg := getSort(sort, direction, "groups") + query.sort += add + query.addGroupBy(agg...) + } + + // Whatever the sorting, always use name/id as a final sort + query.sort += ", COALESCE(groups.name, CAST(groups.id as text)) COLLATE NATURAL_CI ASC" + query.addGroupBy("groups.name", "groups.id") + return nil +} + +func (qb *GroupStore) queryGroups(ctx context.Context, query string, args []interface{}) ([]*models.Group, error) { + const single = false + var ret []*models.Group + if err := groupRepository.queryFunc(ctx, query, args, single, func(r *sqlx.Rows) error { + var f groupRow + if err := r.StructScan(&f); err != nil { + return err + } + + s := f.resolve() + + ret = append(ret, s) + return nil + }); err != nil { + return nil, err + } + + return ret, nil +} + +func (qb *GroupStore) UpdateFrontImage(ctx context.Context, groupID int, frontImage []byte) error { + return qb.UpdateImage(ctx, groupID, groupFrontImageBlobColumn, frontImage) +} + +func (qb *GroupStore) UpdateBackImage(ctx context.Context, groupID int, backImage []byte) error { + return qb.UpdateImage(ctx, groupID, groupBackImageBlobColumn, backImage) +} + +func (qb *GroupStore) destroyImages(ctx context.Context, groupID int) error { + if err := qb.DestroyImage(ctx, groupID, groupFrontImageBlobColumn); err != nil { + return err + } + if err := qb.DestroyImage(ctx, groupID, groupBackImageBlobColumn); err != nil { + return err + } + + return nil +} + +func (qb *GroupStore) GetFrontImage(ctx context.Context, groupID int) ([]byte, error) { + return qb.GetImage(ctx, groupID, groupFrontImageBlobColumn) +} + +func (qb *GroupStore) HasFrontImage(ctx context.Context, groupID int) (bool, error) { + return qb.HasImage(ctx, groupID, groupFrontImageBlobColumn) +} + +func (qb *GroupStore) GetBackImage(ctx context.Context, groupID int) ([]byte, error) { + return qb.GetImage(ctx, groupID, groupBackImageBlobColumn) +} + +func (qb *GroupStore) HasBackImage(ctx context.Context, groupID int) (bool, error) { + return qb.HasImage(ctx, groupID, groupBackImageBlobColumn) +} + +func (qb *GroupStore) FindByPerformerID(ctx context.Context, performerID int) ([]*models.Group, error) { + query := `SELECT DISTINCT groups.* +FROM groups +INNER JOIN groups_scenes ON groups.id = groups_scenes.group_id +INNER JOIN performers_scenes ON performers_scenes.scene_id = groups_scenes.scene_id +WHERE performers_scenes.performer_id = ? +` + args := []interface{}{performerID} + return qb.queryGroups(ctx, query, args) +} + +func (qb *GroupStore) CountByPerformerID(ctx context.Context, performerID int) (int, error) { + query := `SELECT COUNT(DISTINCT groups_scenes.group_id) AS count +FROM groups_scenes +INNER JOIN performers_scenes ON performers_scenes.scene_id = groups_scenes.scene_id +WHERE performers_scenes.performer_id = ? +` + args := []interface{}{performerID} + return groupRepository.runCountQuery(ctx, query, args) +} + +func (qb *GroupStore) FindByStudioID(ctx context.Context, studioID int) ([]*models.Group, error) { + query := `SELECT groups.* +FROM groups +WHERE groups.studio_id = ? +` + args := []interface{}{studioID} + return qb.queryGroups(ctx, query, args) +} + +func (qb *GroupStore) CountByStudioID(ctx context.Context, studioID int) (int, error) { + query := `SELECT COUNT(1) AS count +FROM groups +WHERE groups.studio_id = ? +` + args := []interface{}{studioID} + return groupRepository.runCountQuery(ctx, query, args) +} + +func (qb *GroupStore) GetURLs(ctx context.Context, groupID int) ([]string, error) { + return groupsURLsTableMgr.get(ctx, groupID) +} + +// FindSubGroupIDs returns a list of group IDs where a group in the ids list is a sub-group of the parent group +func (qb *GroupStore) FindSubGroupIDs(ctx context.Context, containingID int, ids []int) ([]int, error) { + /* + SELECT gr.sub_id FROM groups_relations gr + WHERE gr.containing_id = :parentID AND gr.sub_id IN (:ids); + */ + table := groupRelationshipTableMgr.table + q := dialect.From(table).Prepared(true). + Select(table.Col("sub_id")).Where( + table.Col("containing_id").Eq(containingID), + table.Col("sub_id").In(ids), + ) + + const single = false + var ret []int + if err := queryFunc(ctx, q, single, func(r *sqlx.Rows) error { + var id int + if err := r.Scan(&id); err != nil { + return err + } + + ret = append(ret, id) + return nil + }); err != nil { + return nil, err + } + + return ret, nil +} + +// FindInAscestors returns a list of group IDs where a group in the ids list is an ascestor of the ancestor group IDs +func (qb *GroupStore) FindInAncestors(ctx context.Context, ascestorIDs []int, ids []int) ([]int, error) { + /* + WITH RECURSIVE ascestors AS ( + SELECT g.id AS parent_id FROM groups g WHERE g.id IN (:ascestorIDs) + UNION + SELECT gr.containing_id FROM groups_relations gr INNER JOIN ascestors a ON a.parent_id = gr.sub_id + ) + SELECT p.parent_id FROM ascestors p WHERE p.parent_id IN (:ids); + */ + table := qb.table() + const ascestors = "ancestors" + const parentID = "parent_id" + q := dialect.From(ascestors).Prepared(true). + WithRecursive(ascestors, + dialect.From(qb.table()).Select(table.Col(idColumn).As(parentID)). + Where(table.Col(idColumn).In(ascestorIDs)). + Union( + dialect.From(groupRelationsJoinTable).InnerJoin( + goqu.I(ascestors), goqu.On(goqu.I("parent_id").Eq(goqu.I("sub_id"))), + ).Select("containing_id"), + ), + ).Select(parentID).Where(goqu.I(parentID).In(ids)) + + const single = false + var ret []int + if err := queryFunc(ctx, q, single, func(r *sqlx.Rows) error { + var id int + if err := r.Scan(&id); err != nil { + return err + } + + ret = append(ret, id) + return nil + }); err != nil { + return nil, err + } + + return ret, nil +} + +func (qb *GroupStore) sortByOCounter(direction string) string { + // need to sum the o_counter from scenes and images + return " ORDER BY (" + selectGroupOCountSQL + ") " + direction +} diff --git a/pkg/postgres/group_filter.go b/pkg/postgres/group_filter.go new file mode 100644 index 0000000000..eac5c06172 --- /dev/null +++ b/pkg/postgres/group_filter.go @@ -0,0 +1,239 @@ +package postgres + +import ( + "context" + "fmt" + + "github.com/stashapp/stash/pkg/models" + "github.com/stashapp/stash/pkg/utils" +) + +type groupFilterHandler struct { + groupFilter *models.GroupFilterType +} + +func (qb *groupFilterHandler) validate() error { + groupFilter := qb.groupFilter + if groupFilter == nil { + return nil + } + + if err := validateFilterCombination(groupFilter.OperatorFilter); err != nil { + return err + } + + if subFilter := groupFilter.SubFilter(); subFilter != nil { + sqb := &groupFilterHandler{groupFilter: subFilter} + if err := sqb.validate(); err != nil { + return err + } + } + + return nil +} + +func (qb *groupFilterHandler) handle(ctx context.Context, f *filterBuilder) { + groupFilter := qb.groupFilter + if groupFilter == nil { + return + } + + if err := qb.validate(); err != nil { + f.setError(err) + return + } + + sf := groupFilter.SubFilter() + if sf != nil { + sub := &groupFilterHandler{sf} + handleSubFilter(ctx, sub, f, groupFilter.OperatorFilter) + } + + f.handleCriterion(ctx, qb.criterionHandler()) +} + +var groupHierarchyHandler = hierarchicalRelationshipHandler{ + primaryTable: groupTable, + relationTable: groupRelationsTable, + aliasPrefix: groupTable, + parentIDCol: "containing_id", + childIDCol: "sub_id", +} + +func (qb *groupFilterHandler) criterionHandler() criterionHandler { + groupFilter := qb.groupFilter + return compoundHandler{ + stringCriterionHandler(groupFilter.Name, "groups.name"), + stringCriterionHandler(groupFilter.Director, "groups.director"), + stringCriterionHandler(groupFilter.Synopsis, "groups.description"), + intCriterionHandler(groupFilter.Rating100, "groups.rating", nil), + floatIntCriterionHandler(groupFilter.Duration, "groups.duration", nil), + qb.missingCriterionHandler(groupFilter.IsMissing), + qb.urlsCriterionHandler(groupFilter.URL), + studioCriterionHandler(groupTable, groupFilter.Studios), + qb.performersCriterionHandler(groupFilter.Performers), + qb.tagsCriterionHandler(groupFilter.Tags), + qb.tagCountCriterionHandler(groupFilter.TagCount), + qb.groupOCounterCriterionHandler(groupFilter.OCounter), + &dateCriterionHandler{groupFilter.Date, "groups.date", nil}, + groupHierarchyHandler.ParentsCriterionHandler(groupFilter.ContainingGroups), + groupHierarchyHandler.ChildrenCriterionHandler(groupFilter.SubGroups), + groupHierarchyHandler.ParentCountCriterionHandler(groupFilter.ContainingGroupCount), + groupHierarchyHandler.ChildCountCriterionHandler(groupFilter.SubGroupCount), + ×tampCriterionHandler{groupFilter.CreatedAt, "groups.created_at", nil}, + ×tampCriterionHandler{groupFilter.UpdatedAt, "groups.updated_at", nil}, + + &relatedFilterHandler{ + relatedIDCol: "groups_scenes.scene_id", + relatedRepo: sceneRepository.repository, + relatedHandler: &sceneFilterHandler{groupFilter.ScenesFilter}, + joinFn: func(f *filterBuilder) { + groupRepository.scenes.innerJoin(f, "", "groups.id") + }, + }, + + &relatedFilterHandler{ + relatedIDCol: "groups.studio_id", + relatedRepo: studioRepository.repository, + relatedHandler: &studioFilterHandler{groupFilter.StudiosFilter}, + }, + } +} + +func (qb *groupFilterHandler) missingCriterionHandler(isMissing *string) criterionHandlerFunc { + return func(ctx context.Context, f *filterBuilder) { + if isMissing != nil && *isMissing != "" { + switch *isMissing { + case "front_image": + f.addWhere("groups.front_image_blob IS NULL") + case "back_image": + f.addWhere("groups.back_image_blob IS NULL") + case "scenes": + f.addLeftJoin("groups_scenes", "", "groups_scenes.group_id = groups.id") + f.addWhere("groups_scenes.scene_id IS NULL") + default: + f.addWhere("(groups." + *isMissing + " IS NULL OR TRIM(CAST(groups." + *isMissing + " AS TEXT)) = '')") + } + } + } +} + +func (qb *groupFilterHandler) urlsCriterionHandler(url *models.StringCriterionInput) criterionHandlerFunc { + h := stringListCriterionHandlerBuilder{ + primaryTable: groupTable, + primaryFK: groupIDColumn, + joinTable: groupURLsTable, + stringColumn: groupURLColumn, + addJoinTable: func(f *filterBuilder) { + groupsURLsTableMgr.join(f, "", "groups.id") + }, + } + + return h.handler(url) +} + +func (qb *groupFilterHandler) performersCriterionHandler(performers *models.MultiCriterionInput) criterionHandlerFunc { + return func(ctx context.Context, f *filterBuilder) { + if performers != nil { + if performers.Modifier == models.CriterionModifierIsNull || performers.Modifier == models.CriterionModifierNotNull { + var notClause string + if performers.Modifier == models.CriterionModifierNotNull { + notClause = "NOT" + } + + f.addLeftJoin("groups_scenes", "", "groups.id = groups_scenes.group_id") + f.addLeftJoin("performers_scenes", "", "groups_scenes.scene_id = performers_scenes.scene_id") + + f.addWhere(fmt.Sprintf("performers_scenes.performer_id IS %s NULL", notClause)) + return + } + + if len(performers.Value) == 0 { + return + } + + var args []interface{} + for _, arg := range performers.Value { + args = append(args, arg) + } + + // Hack, can't apply args to join, nor inner join on a left join, so use CTE instead + f.addWith(`groups_performers AS ( + SELECT groups_scenes.group_id, performers_scenes.performer_id + FROM groups_scenes + INNER JOIN performers_scenes ON groups_scenes.scene_id = performers_scenes.scene_id + WHERE performers_scenes.performer_id IN`+getInBinding(len(performers.Value))+` + )`, args...) + f.addLeftJoin("groups_performers", "", "groups.id = groups_performers.group_id") + + switch performers.Modifier { + case models.CriterionModifierIncludes: + f.addWhere("groups_performers.performer_id IS NOT NULL") + case models.CriterionModifierIncludesAll: + f.addWhere("groups_performers.performer_id IS NOT NULL") + f.addHaving("COUNT(DISTINCT groups_performers.performer_id) = ?", len(performers.Value)) + case models.CriterionModifierExcludes: + f.addWhere("groups_performers.performer_id IS NULL") + } + } + } +} + +func (qb *groupFilterHandler) tagsCriterionHandler(tags *models.HierarchicalMultiCriterionInput) criterionHandlerFunc { + h := joinedHierarchicalMultiCriterionHandlerBuilder{ + primaryTable: groupTable, + foreignTable: tagTable, + foreignFK: "tag_id", + + relationsTable: "tags_relations", + joinAs: "group_tag", + joinTable: groupsTagsTable, + primaryFK: groupIDColumn, + } + + return h.handler(tags) +} + +func (qb *groupFilterHandler) tagCountCriterionHandler(count *models.IntCriterionInput) criterionHandlerFunc { + h := countCriterionHandlerBuilder{ + primaryTable: groupTable, + joinTable: groupsTagsTable, + primaryFK: groupIDColumn, + } + + return h.handler(count) +} + +// used for sorting and filtering on group o-count +var selectGroupOCountSQL = utils.StrFormat( + "SELECT SUM(o_counter) "+ + "FROM ("+ + "SELECT COUNT({scenes_o_dates}.{o_date}) as o_counter from {groups_scenes} s "+ + "LEFT JOIN {scenes} ON {scenes}.id = s.{scene_id} "+ + "LEFT JOIN {scenes_o_dates} ON {scenes_o_dates}.{scene_id} = {scenes}.id "+ + "WHERE s.{group_id} = {group}.id "+ + ")", + map[string]interface{}{ + "group": groupTable, + "group_id": groupIDColumn, + "groups_scenes": groupsScenesTable, + "scenes": sceneTable, + "scene_id": sceneIDColumn, + "scenes_o_dates": scenesODatesTable, + "o_date": sceneODateColumn, + }, +) + +func (qb *groupFilterHandler) groupOCounterCriterionHandler(count *models.IntCriterionInput) criterionHandlerFunc { + return func(ctx context.Context, f *filterBuilder) { + if count == nil { + return + } + + lhs := "(" + selectGroupOCountSQL + ")" + clause, args := getIntCriterionWhereClause(lhs, *count) + + f.addWhere(clause, args...) + } + +} diff --git a/pkg/postgres/group_relationships.go b/pkg/postgres/group_relationships.go new file mode 100644 index 0000000000..025949e0b4 --- /dev/null +++ b/pkg/postgres/group_relationships.go @@ -0,0 +1,457 @@ +package postgres + +import ( + "context" + "fmt" + + "github.com/doug-martin/goqu/v9" + "github.com/doug-martin/goqu/v9/exp" + "github.com/jmoiron/sqlx" + "github.com/stashapp/stash/pkg/models" + "gopkg.in/guregu/null.v4" + "gopkg.in/guregu/null.v4/zero" +) + +type groupRelationshipRow struct { + ContainingID int `db:"containing_id"` + SubID int `db:"sub_id"` + OrderIndex int `db:"order_index"` + Description zero.String `db:"description"` +} + +func (r groupRelationshipRow) resolve(useContainingID bool) models.GroupIDDescription { + id := r.ContainingID + if !useContainingID { + id = r.SubID + } + + return models.GroupIDDescription{ + GroupID: id, + Description: r.Description.String, + } +} + +type groupRelationshipStore struct { + table *table +} + +func (s *groupRelationshipStore) GetContainingGroupDescriptions(ctx context.Context, id int) ([]models.GroupIDDescription, error) { + const idIsContaining = false + return s.getGroupRelationships(ctx, id, idIsContaining) +} + +func (s *groupRelationshipStore) GetSubGroupDescriptions(ctx context.Context, id int) ([]models.GroupIDDescription, error) { + const idIsContaining = true + return s.getGroupRelationships(ctx, id, idIsContaining) +} + +func (s *groupRelationshipStore) getGroupRelationships(ctx context.Context, id int, idIsContaining bool) ([]models.GroupIDDescription, error) { + col := "containing_id" + if !idIsContaining { + col = "sub_id" + } + + table := s.table.table + q := dialect.Select(table.All()). + From(table). + Where(table.Col(col).Eq(id)). + Order(table.Col("order_index").Asc()) + + const single = false + var ret []models.GroupIDDescription + if err := queryFunc(ctx, q, single, func(rows *sqlx.Rows) error { + var row groupRelationshipRow + if err := rows.StructScan(&row); err != nil { + return err + } + + ret = append(ret, row.resolve(!idIsContaining)) + + return nil + }); err != nil { + return nil, fmt.Errorf("getting group relationships from %s: %w", table.GetTable(), err) + } + + return ret, nil +} + +// getMaxOrderIndex gets the maximum order index for the containing group with the given id +func (s *groupRelationshipStore) getMaxOrderIndex(ctx context.Context, containingID int) (int, error) { + idColumn := s.table.table.Col("containing_id") + + q := dialect.Select(goqu.MAX("order_index")). + From(s.table.table). + Where(idColumn.Eq(containingID)) + + var maxOrderIndex zero.Int + if err := querySimple(ctx, q, &maxOrderIndex); err != nil { + return 0, fmt.Errorf("getting max order index: %w", err) + } + + return int(maxOrderIndex.Int64), nil +} + +// createRelationships creates relationships between a group and other groups. +// If idIsContaining is true, the provided id is the containing group. +func (s *groupRelationshipStore) createRelationships(ctx context.Context, id int, d models.RelatedGroupDescriptions, idIsContaining bool) error { + if d.Loaded() { + for i, v := range d.List() { + orderIndex := i + 1 + + r := groupRelationshipRow{ + ContainingID: id, + SubID: v.GroupID, + OrderIndex: orderIndex, + Description: zero.StringFrom(v.Description), + } + + if !idIsContaining { + // get the max order index of the containing groups sub groups + containingID := v.GroupID + maxOrderIndex, err := s.getMaxOrderIndex(ctx, containingID) + if err != nil { + return err + } + + r.ContainingID = v.GroupID + r.SubID = id + r.OrderIndex = maxOrderIndex + 1 + } + + _, err := s.table.insert(ctx, r) + if err != nil { + return fmt.Errorf("inserting into %s: %w", s.table.table.GetTable(), err) + } + } + + return nil + } + + return nil +} + +// createRelationships creates relationships between a group and other groups. +// If idIsContaining is true, the provided id is the containing group. +func (s *groupRelationshipStore) createContainingRelationships(ctx context.Context, id int, d models.RelatedGroupDescriptions) error { + const idIsContaining = false + return s.createRelationships(ctx, id, d, idIsContaining) +} + +// createRelationships creates relationships between a group and other groups. +// If idIsContaining is true, the provided id is the containing group. +func (s *groupRelationshipStore) createSubRelationships(ctx context.Context, id int, d models.RelatedGroupDescriptions) error { + const idIsContaining = true + return s.createRelationships(ctx, id, d, idIsContaining) +} + +func (s *groupRelationshipStore) replaceRelationships(ctx context.Context, id int, d models.RelatedGroupDescriptions, idIsContaining bool) error { + // always destroy the existing relationships even if the new list is empty + if err := s.destroyAllJoins(ctx, id, idIsContaining); err != nil { + return err + } + + return s.createRelationships(ctx, id, d, idIsContaining) +} + +func (s *groupRelationshipStore) replaceContainingRelationships(ctx context.Context, id int, d models.RelatedGroupDescriptions) error { + const idIsContaining = false + return s.replaceRelationships(ctx, id, d, idIsContaining) +} + +func (s *groupRelationshipStore) replaceSubRelationships(ctx context.Context, id int, d models.RelatedGroupDescriptions) error { + const idIsContaining = true + return s.replaceRelationships(ctx, id, d, idIsContaining) +} + +func (s *groupRelationshipStore) modifyRelationships(ctx context.Context, id int, v *models.UpdateGroupDescriptions, idIsContaining bool) error { + if v == nil { + return nil + } + + switch v.Mode { + case models.RelationshipUpdateModeSet: + return s.replaceJoins(ctx, id, *v, idIsContaining) + case models.RelationshipUpdateModeAdd: + return s.addJoins(ctx, id, v.Groups, idIsContaining) + case models.RelationshipUpdateModeRemove: + toRemove := make([]int, len(v.Groups)) + for i, vv := range v.Groups { + toRemove[i] = vv.GroupID + } + return s.destroyJoins(ctx, id, toRemove, idIsContaining) + } + + return nil +} + +func (s *groupRelationshipStore) modifyContainingRelationships(ctx context.Context, id int, v *models.UpdateGroupDescriptions) error { + const idIsContaining = false + return s.modifyRelationships(ctx, id, v, idIsContaining) +} + +func (s *groupRelationshipStore) modifySubRelationships(ctx context.Context, id int, v *models.UpdateGroupDescriptions) error { + const idIsContaining = true + return s.modifyRelationships(ctx, id, v, idIsContaining) +} + +func (s *groupRelationshipStore) addJoins(ctx context.Context, id int, groups []models.GroupIDDescription, idIsContaining bool) error { + // if we're adding to a containing group, get the max order index first + var maxOrderIndex int + if idIsContaining { + var err error + maxOrderIndex, err = s.getMaxOrderIndex(ctx, id) + if err != nil { + return err + } + } + + for i, vv := range groups { + r := groupRelationshipRow{ + Description: zero.StringFrom(vv.Description), + } + + if idIsContaining { + r.ContainingID = id + r.SubID = vv.GroupID + r.OrderIndex = maxOrderIndex + (i + 1) + } else { + // get the max order index of the containing groups sub groups + containingMaxOrderIndex, err := s.getMaxOrderIndex(ctx, vv.GroupID) + if err != nil { + return err + } + + r.ContainingID = vv.GroupID + r.SubID = id + r.OrderIndex = containingMaxOrderIndex + 1 + } + + _, err := s.table.insert(ctx, r) + if err != nil { + return fmt.Errorf("inserting into %s: %w", s.table.table.GetTable(), err) + } + } + + return nil +} + +func (s *groupRelationshipStore) destroyAllJoins(ctx context.Context, id int, idIsContaining bool) error { + table := s.table.table + idColumn := table.Col("containing_id") + if !idIsContaining { + idColumn = table.Col("sub_id") + } + + q := dialect.Delete(table).Where(idColumn.Eq(id)) + + if _, err := exec(ctx, q); err != nil { + return fmt.Errorf("destroying %s: %w", table.GetTable(), err) + } + + return nil +} + +func (s *groupRelationshipStore) replaceJoins(ctx context.Context, id int, v models.UpdateGroupDescriptions, idIsContaining bool) error { + if err := s.destroyAllJoins(ctx, id, idIsContaining); err != nil { + return err + } + + // convert to RelatedGroupDescriptions + rgd := models.NewRelatedGroupDescriptions(v.Groups) + return s.createRelationships(ctx, id, rgd, idIsContaining) +} + +func (s *groupRelationshipStore) destroyJoins(ctx context.Context, id int, toRemove []int, idIsContaining bool) error { + table := s.table.table + idColumn := table.Col("containing_id") + fkColumn := table.Col("sub_id") + if !idIsContaining { + idColumn = table.Col("sub_id") + fkColumn = table.Col("containing_id") + } + + q := dialect.Delete(table).Where(idColumn.Eq(id), fkColumn.In(toRemove)) + + if _, err := exec(ctx, q); err != nil { + return fmt.Errorf("destroying %s: %w", table.GetTable(), err) + } + + return nil +} + +func (s *groupRelationshipStore) getOrderIndexOfSubGroup(ctx context.Context, containingGroupID int, subGroupID int) (int, error) { + table := s.table.table + q := dialect.Select("order_index"). + From(table). + Where( + table.Col("containing_id").Eq(containingGroupID), + table.Col("sub_id").Eq(subGroupID), + ) + + var orderIndex null.Int + if err := querySimple(ctx, q, &orderIndex); err != nil { + return 0, fmt.Errorf("getting order index: %w", err) + } + + if !orderIndex.Valid { + return 0, fmt.Errorf("sub-group %d not found in containing group %d", subGroupID, containingGroupID) + } + + return int(orderIndex.Int64), nil +} + +func (s *groupRelationshipStore) getGroupIDAtOrderIndex(ctx context.Context, containingGroupID int, orderIndex int) (*int, error) { + table := s.table.table + q := dialect.Select(table.Col("sub_id")).From(table).Where( + table.Col("containing_id").Eq(containingGroupID), + table.Col("order_index").Eq(orderIndex), + ) + + var ret null.Int + if err := querySimple(ctx, q, &ret); err != nil { + return nil, fmt.Errorf("getting sub id for order index: %w", err) + } + + if !ret.Valid { + return nil, nil + } + + intRet := int(ret.Int64) + return &intRet, nil +} + +func (s *groupRelationshipStore) getOrderIndexAfterOrderIndex(ctx context.Context, containingGroupID int, orderIndex int) (int, error) { + table := s.table.table + q := dialect.Select(goqu.MIN("order_index")).From(table).Where( + table.Col("containing_id").Eq(containingGroupID), + table.Col("order_index").Gt(orderIndex), + ) + + var ret null.Int + if err := querySimple(ctx, q, &ret); err != nil { + return 0, fmt.Errorf("getting order index: %w", err) + } + + if !ret.Valid { + return orderIndex + 1, nil + } + + return int(ret.Int64), nil +} + +// incrementOrderIndexes increments the order_index value of all sub-groups in the containing group at or after the given index +func (s *groupRelationshipStore) incrementOrderIndexes(ctx context.Context, groupID int, indexBefore int) error { + table := s.table.table + + // WORKAROUND - sqlite won't allow incrementing the value directly since it causes a + // unique constraint violation. + // Instead, we first set the order index to a negative value temporarily + // see https://stackoverflow.com/a/7703239/695786 + q := dialect.Update(table).Set(exp.Record{ + "order_index": goqu.L("-order_index"), + }).Where( + table.Col("containing_id").Eq(groupID), + table.Col("order_index").Gte(indexBefore), + ) + + if _, err := exec(ctx, q); err != nil { + return fmt.Errorf("updating %s: %w", table.GetTable(), err) + } + + q = dialect.Update(table).Set(exp.Record{ + "order_index": goqu.L("1-order_index"), + }).Where( + table.Col("containing_id").Eq(groupID), + table.Col("order_index").Lt(0), + ) + + if _, err := exec(ctx, q); err != nil { + return fmt.Errorf("updating %s: %w", table.GetTable(), err) + } + + return nil +} + +func (s *groupRelationshipStore) reorderSubGroup(ctx context.Context, groupID int, subGroupID int, insertPointID int, insertAfter bool) error { + insertPointIndex, err := s.getOrderIndexOfSubGroup(ctx, groupID, insertPointID) + if err != nil { + return err + } + + // if we're setting before + if insertAfter { + insertPointIndex, err = s.getOrderIndexAfterOrderIndex(ctx, groupID, insertPointIndex) + if err != nil { + return err + } + } + + // increment the order index of all sub-groups after and including the insertion point + if err := s.incrementOrderIndexes(ctx, groupID, int(insertPointIndex)); err != nil { + return err + } + + // set the order index of the sub-group to the insertion point + table := s.table.table + q := dialect.Update(table).Set(exp.Record{ + "order_index": insertPointIndex, + }).Where( + table.Col("containing_id").Eq(groupID), + table.Col("sub_id").Eq(subGroupID), + ) + + if _, err := exec(ctx, q); err != nil { + return fmt.Errorf("updating %s: %w", table.GetTable(), err) + } + + return nil +} + +func (s *groupRelationshipStore) AddSubGroups(ctx context.Context, groupID int, subGroups []models.GroupIDDescription, insertIndex *int) error { + const idIsContaining = true + + if err := s.addJoins(ctx, groupID, subGroups, idIsContaining); err != nil { + return err + } + + ids := make([]int, len(subGroups)) + for i, v := range subGroups { + ids[i] = v.GroupID + } + + if insertIndex != nil { + // get the id of the sub-group at the insert index + insertPointID, err := s.getGroupIDAtOrderIndex(ctx, groupID, *insertIndex) + if err != nil { + return err + } + + if insertPointID == nil { + // if the insert index is out of bounds, just assume adding to the end + return nil + } + + // reorder the sub-groups + const insertAfter = false + if err := s.ReorderSubGroups(ctx, groupID, ids, *insertPointID, insertAfter); err != nil { + return err + } + } + + return nil +} + +func (s *groupRelationshipStore) RemoveSubGroups(ctx context.Context, groupID int, subGroupIDs []int) error { + const idIsContaining = true + return s.destroyJoins(ctx, groupID, subGroupIDs, idIsContaining) +} + +func (s *groupRelationshipStore) ReorderSubGroups(ctx context.Context, groupID int, subGroupIDs []int, insertPointID int, insertAfter bool) error { + for _, id := range subGroupIDs { + if err := s.reorderSubGroup(ctx, groupID, id, insertPointID, insertAfter); err != nil { + return err + } + } + + return nil +} diff --git a/pkg/postgres/history.go b/pkg/postgres/history.go new file mode 100644 index 0000000000..bb15d770b2 --- /dev/null +++ b/pkg/postgres/history.go @@ -0,0 +1,95 @@ +package postgres + +import ( + "context" + "time" +) + +type viewDateManager struct { + tableMgr *viewHistoryTable +} + +func (qb *viewDateManager) GetViewDates(ctx context.Context, id int) ([]time.Time, error) { + return qb.tableMgr.getDates(ctx, id) +} + +func (qb *viewDateManager) GetManyViewDates(ctx context.Context, ids []int) ([][]time.Time, error) { + return qb.tableMgr.getManyDates(ctx, ids) +} + +func (qb *viewDateManager) CountViews(ctx context.Context, id int) (int, error) { + return qb.tableMgr.getCount(ctx, id) +} + +func (qb *viewDateManager) GetManyViewCount(ctx context.Context, ids []int) ([]int, error) { + return qb.tableMgr.getManyCount(ctx, ids) +} + +func (qb *viewDateManager) CountAllViews(ctx context.Context) (int, error) { + return qb.tableMgr.getAllCount(ctx) +} + +func (qb *viewDateManager) CountUniqueViews(ctx context.Context) (int, error) { + return qb.tableMgr.getUniqueCount(ctx) +} + +func (qb *viewDateManager) LastView(ctx context.Context, id int) (*time.Time, error) { + return qb.tableMgr.getLastDate(ctx, id) +} + +func (qb *viewDateManager) GetManyLastViewed(ctx context.Context, ids []int) ([]*time.Time, error) { + return qb.tableMgr.getManyLastDate(ctx, ids) + +} + +func (qb *viewDateManager) AddViews(ctx context.Context, id int, dates []time.Time) ([]time.Time, error) { + return qb.tableMgr.addDates(ctx, id, dates) +} + +func (qb *viewDateManager) DeleteViews(ctx context.Context, id int, dates []time.Time) ([]time.Time, error) { + return qb.tableMgr.deleteDates(ctx, id, dates) +} + +func (qb *viewDateManager) DeleteAllViews(ctx context.Context, id int) (int, error) { + return qb.tableMgr.deleteAllDates(ctx, id) +} + +type oDateManager struct { + tableMgr *viewHistoryTable +} + +func (qb *oDateManager) GetODates(ctx context.Context, id int) ([]time.Time, error) { + return qb.tableMgr.getDates(ctx, id) +} + +func (qb *oDateManager) GetManyODates(ctx context.Context, ids []int) ([][]time.Time, error) { + return qb.tableMgr.getManyDates(ctx, ids) +} + +func (qb *oDateManager) GetOCount(ctx context.Context, id int) (int, error) { + return qb.tableMgr.getCount(ctx, id) +} + +func (qb *oDateManager) GetManyOCount(ctx context.Context, ids []int) ([]int, error) { + return qb.tableMgr.getManyCount(ctx, ids) +} + +func (qb *oDateManager) GetAllOCount(ctx context.Context) (int, error) { + return qb.tableMgr.getAllCount(ctx) +} + +func (qb *oDateManager) GetUniqueOCount(ctx context.Context) (int, error) { + return qb.tableMgr.getUniqueCount(ctx) +} + +func (qb *oDateManager) AddO(ctx context.Context, id int, dates []time.Time) ([]time.Time, error) { + return qb.tableMgr.addDates(ctx, id, dates) +} + +func (qb *oDateManager) DeleteO(ctx context.Context, id int, dates []time.Time) ([]time.Time, error) { + return qb.tableMgr.deleteDates(ctx, id, dates) +} + +func (qb *oDateManager) ResetO(ctx context.Context, id int) (int, error) { + return qb.tableMgr.deleteAllDates(ctx, id) +} diff --git a/pkg/postgres/image.go b/pkg/postgres/image.go new file mode 100644 index 0000000000..a34544f235 --- /dev/null +++ b/pkg/postgres/image.go @@ -0,0 +1,1093 @@ +package postgres + +import ( + "context" + "database/sql" + "errors" + "fmt" + "path/filepath" + "slices" + + "github.com/jmoiron/sqlx" + "github.com/stashapp/stash/pkg/models" + "github.com/stashapp/stash/pkg/sliceutil" + "gopkg.in/guregu/null.v4" + "gopkg.in/guregu/null.v4/zero" + + "github.com/doug-martin/goqu/v9" + "github.com/doug-martin/goqu/v9/exp" +) + +const imageTable = "images" + +const ( + imageIDColumn = "image_id" + performersImagesTable = "performers_images" + imagesTagsTable = "images_tags" + imagesFilesTable = "images_files" + imagesURLsTable = "image_urls" + imageURLColumn = "url" +) + +type imageRow struct { + ID int `db:"id" goqu:"skipinsert"` + Title zero.String `db:"title"` + Code zero.String `db:"code"` + // expressed as 1-100 + Rating null.Int `db:"rating"` + Date NullDate `db:"date"` + DatePrecision null.Int `db:"date_precision"` + Details zero.String `db:"details"` + Photographer zero.String `db:"photographer"` + Organized bool `db:"organized"` + OCounter int `db:"o_counter"` + StudioID null.Int `db:"studio_id,omitempty"` + CreatedAt Timestamp `db:"created_at"` + UpdatedAt Timestamp `db:"updated_at"` +} + +func (r *imageRow) fromImage(i models.Image) { + r.ID = i.ID + r.Title = zero.StringFrom(i.Title) + r.Code = zero.StringFrom(i.Code) + r.Rating = intFromPtr(i.Rating) + r.Date = NullDateFromDatePtr(i.Date) + r.DatePrecision = datePrecisionFromDatePtr(i.Date) + r.Details = zero.StringFrom(i.Details) + r.Photographer = zero.StringFrom(i.Photographer) + r.Organized = i.Organized + r.OCounter = i.OCounter + r.StudioID = intFromPtr(i.StudioID) + r.CreatedAt = Timestamp{Timestamp: i.CreatedAt} + r.UpdatedAt = Timestamp{Timestamp: i.UpdatedAt} +} + +type imageQueryRow struct { + imageRow + PrimaryFileID null.Int `db:"primary_file_id"` + PrimaryFileFolderPath zero.String `db:"primary_file_folder_path"` + PrimaryFileBasename zero.String `db:"primary_file_basename"` + PrimaryFileChecksum zero.String `db:"primary_file_checksum"` +} + +func (r *imageQueryRow) resolve() *models.Image { + ret := &models.Image{ + ID: r.ID, + Title: r.Title.String, + Code: r.Code.String, + Rating: nullIntPtr(r.Rating), + Date: r.Date.DatePtr(r.DatePrecision), + Details: r.Details.String, + Photographer: r.Photographer.String, + Organized: r.Organized, + OCounter: r.OCounter, + StudioID: nullIntPtr(r.StudioID), + + PrimaryFileID: nullIntFileIDPtr(r.PrimaryFileID), + Checksum: r.PrimaryFileChecksum.String, + + CreatedAt: r.CreatedAt.Timestamp.UTC(), + UpdatedAt: r.UpdatedAt.Timestamp.UTC(), + } + + if r.PrimaryFileFolderPath.Valid && r.PrimaryFileBasename.Valid { + ret.Path = filepath.Join(r.PrimaryFileFolderPath.String, r.PrimaryFileBasename.String) + } + + return ret +} + +type imageRowRecord struct { + updateRecord +} + +func (r *imageRowRecord) fromPartial(i models.ImagePartial) { + r.setNullString("title", i.Title) + r.setNullString("code", i.Code) + r.setNullInt("rating", i.Rating) + r.setNullDate("date", "date_precision", i.Date) + r.setNullString("details", i.Details) + r.setNullString("photographer", i.Photographer) + r.setBool("organized", i.Organized) + r.setInt("o_counter", i.OCounter) + r.setNullInt("studio_id", i.StudioID) + r.setTimestamp("created_at", i.CreatedAt) + r.setTimestamp("updated_at", i.UpdatedAt) +} + +type imageRepositoryType struct { + repository + performers joinRepository + galleries joinRepository + tags joinRepository + files filesRepository +} + +func (r *imageRepositoryType) addImagesFilesTable(f *filterBuilder) { + f.addLeftJoin(imagesFilesTable, "", "images_files.image_id = images.id AND images_files.\"primary\" = true") +} + +func (r *imageRepositoryType) addFilesTable(f *filterBuilder) { + r.addImagesFilesTable(f) + f.addLeftJoin(fileTable, "", "images_files.file_id = files.id") +} + +func (r *imageRepositoryType) addFoldersTable(f *filterBuilder) { + r.addFilesTable(f) + f.addLeftJoin(folderTable, "", "files.parent_folder_id = folders.id") +} + +func (r *imageRepositoryType) addImageFilesTable(f *filterBuilder) { + r.addImagesFilesTable(f) + f.addLeftJoin(imageFileTable, "", "image_files.file_id = images_files.file_id") +} + +var ( + imageRepository = imageRepositoryType{ + repository: repository{ + tableName: imageTable, + idColumn: idColumn, + }, + + performers: joinRepository{ + repository: repository{ + tableName: performersImagesTable, + idColumn: imageIDColumn, + }, + fkColumn: performerIDColumn, + }, + + galleries: joinRepository{ + repository: repository{ + tableName: galleriesImagesTable, + idColumn: imageIDColumn, + }, + fkColumn: galleryIDColumn, + }, + + files: filesRepository{ + repository: repository{ + tableName: imagesFilesTable, + idColumn: imageIDColumn, + }, + }, + + tags: joinRepository{ + repository: repository{ + tableName: imagesTagsTable, + idColumn: imageIDColumn, + }, + fkColumn: tagIDColumn, + foreignTable: tagTable, + orderBy: tagTableSortSQL, + }, + } +) + +type ImageStore struct { + tableMgr *table + oCounterManager + + repo *storeRepository +} + +func NewImageStore(r *storeRepository) *ImageStore { + return &ImageStore{ + tableMgr: imageTableMgr, + oCounterManager: oCounterManager{imageTableMgr}, + repo: r, + } +} + +func (qb *ImageStore) table() exp.IdentifierExpression { + return qb.tableMgr.table +} + +func (qb *ImageStore) selectDataset() *goqu.SelectDataset { + table := qb.table() + files := fileTableMgr.table + folders := folderTableMgr.table + checksum := fingerprintTableMgr.table + + return dialect.From(table).LeftJoin( + imagesFilesJoinTable, + goqu.On( + imagesFilesJoinTable.Col(imageIDColumn).Eq(table.Col(idColumn)), + imagesFilesJoinTable.Col("primary").IsTrue(), + ), + ).LeftJoin( + files, + goqu.On(files.Col(idColumn).Eq(imagesFilesJoinTable.Col(fileIDColumn))), + ).LeftJoin( + folders, + goqu.On(folders.Col(idColumn).Eq(files.Col("parent_folder_id"))), + ).LeftJoin( + checksum, + goqu.On( + checksum.Col(fileIDColumn).Eq(imagesFilesJoinTable.Col(fileIDColumn)), + checksum.Col("type").Eq(models.FingerprintTypeMD5), + ), + ).Select( + qb.table().All(), + imagesFilesJoinTable.Col(fileIDColumn).As("primary_file_id"), + folders.Col("path").As("primary_file_folder_path"), + files.Col("basename").As("primary_file_basename"), + checksum.Col("fingerprint").As("primary_file_checksum"), + ) +} + +func (qb *ImageStore) Create(ctx context.Context, newObject *models.Image, fileIDs []models.FileID) error { + var r imageRow + r.fromImage(*newObject) + + id, err := qb.tableMgr.insertID(ctx, r) + if err != nil { + return err + } + + if len(fileIDs) > 0 { + const firstPrimary = true + if err := imagesFilesTableMgr.insertJoins(ctx, id, firstPrimary, fileIDs); err != nil { + return err + } + } + + if newObject.URLs.Loaded() { + const startPos = 0 + if err := imagesURLsTableMgr.insertJoins(ctx, id, startPos, newObject.URLs.List()); err != nil { + return err + } + } + + if newObject.PerformerIDs.Loaded() { + if err := imagesPerformersTableMgr.insertJoins(ctx, id, newObject.PerformerIDs.List()); err != nil { + return err + } + } + if newObject.TagIDs.Loaded() { + if err := imagesTagsTableMgr.insertJoins(ctx, id, newObject.TagIDs.List()); err != nil { + return err + } + } + + if newObject.GalleryIDs.Loaded() { + if err := imageGalleriesTableMgr.insertJoins(ctx, id, newObject.GalleryIDs.List()); err != nil { + return err + } + } + + updated, err := qb.find(ctx, id) + if err != nil { + return fmt.Errorf("finding after create: %w", err) + } + + *newObject = *updated + + return nil +} + +func (qb *ImageStore) UpdatePartial(ctx context.Context, id int, partial models.ImagePartial) (*models.Image, error) { + r := imageRowRecord{ + updateRecord{ + Record: make(exp.Record), + }, + } + + r.fromPartial(partial) + + if len(r.Record) > 0 { + if err := qb.tableMgr.updateByID(ctx, id, r.Record); err != nil { + return nil, err + } + } + + if partial.GalleryIDs != nil { + if err := imageGalleriesTableMgr.modifyJoins(ctx, id, partial.GalleryIDs.IDs, partial.GalleryIDs.Mode); err != nil { + return nil, err + } + } + + if partial.URLs != nil { + if err := imagesURLsTableMgr.modifyJoins(ctx, id, partial.URLs.Values, partial.URLs.Mode); err != nil { + return nil, err + } + } + if partial.PerformerIDs != nil { + if err := imagesPerformersTableMgr.modifyJoins(ctx, id, partial.PerformerIDs.IDs, partial.PerformerIDs.Mode); err != nil { + return nil, err + } + } + if partial.TagIDs != nil { + if err := imagesTagsTableMgr.modifyJoins(ctx, id, partial.TagIDs.IDs, partial.TagIDs.Mode); err != nil { + return nil, err + } + } + + if partial.PrimaryFileID != nil { + if err := imagesFilesTableMgr.setPrimary(ctx, id, *partial.PrimaryFileID); err != nil { + return nil, err + } + } + + return qb.find(ctx, id) +} + +func (qb *ImageStore) Update(ctx context.Context, updatedObject *models.Image) error { + var r imageRow + r.fromImage(*updatedObject) + + if err := qb.tableMgr.updateByID(ctx, updatedObject.ID, r); err != nil { + return err + } + + if updatedObject.URLs.Loaded() { + if err := imagesURLsTableMgr.replaceJoins(ctx, updatedObject.ID, updatedObject.URLs.List()); err != nil { + return err + } + } + + if updatedObject.PerformerIDs.Loaded() { + if err := imagesPerformersTableMgr.replaceJoins(ctx, updatedObject.ID, updatedObject.PerformerIDs.List()); err != nil { + return err + } + } + + if updatedObject.TagIDs.Loaded() { + if err := imagesTagsTableMgr.replaceJoins(ctx, updatedObject.ID, updatedObject.TagIDs.List()); err != nil { + return err + } + } + + if updatedObject.GalleryIDs.Loaded() { + if err := imageGalleriesTableMgr.replaceJoins(ctx, updatedObject.ID, updatedObject.GalleryIDs.List()); err != nil { + return err + } + } + + if updatedObject.Files.Loaded() { + fileIDs := make([]models.FileID, len(updatedObject.Files.List())) + for i, f := range updatedObject.Files.List() { + fileIDs[i] = f.Base().ID + } + + if err := imagesFilesTableMgr.replaceJoins(ctx, updatedObject.ID, fileIDs); err != nil { + return err + } + } + return nil +} + +func (qb *ImageStore) Destroy(ctx context.Context, id int) error { + return qb.tableMgr.destroyExisting(ctx, []int{id}) +} + +// returns nil, nil if not found +func (qb *ImageStore) Find(ctx context.Context, id int) (*models.Image, error) { + ret, err := qb.find(ctx, id) + if errors.Is(err, sql.ErrNoRows) { + return nil, nil + } + return ret, err +} + +func (qb *ImageStore) FindMany(ctx context.Context, ids []int) ([]*models.Image, error) { + images := make([]*models.Image, len(ids)) + + if len(ids) == 0 { + return images, nil + } + + if err := batchExec(ids, defaultBatchSize, func(batch []int) error { + q := qb.selectDataset().Prepared(true).Where(qb.table().Col(idColumn).In(batch)) + unsorted, err := qb.getMany(ctx, q) + if err != nil { + return err + } + + for _, s := range unsorted { + i := slices.Index(ids, s.ID) + images[i] = s + } + + return nil + }); err != nil { + return nil, err + } + + for i := range images { + if images[i] == nil { + return nil, fmt.Errorf("image with id %d not found", ids[i]) + } + } + + return images, nil +} + +// returns nil, sql.ErrNoRows if not found +func (qb *ImageStore) find(ctx context.Context, id int) (*models.Image, error) { + q := qb.selectDataset().Where(qb.tableMgr.byID(id)) + + ret, err := qb.get(ctx, q) + if err != nil { + return nil, err + } + + return ret, nil +} + +func (qb *ImageStore) findBySubquery(ctx context.Context, sq *goqu.SelectDataset) ([]*models.Image, error) { + table := qb.table() + + q := qb.selectDataset().Prepared(true).Where( + table.Col(idColumn).Eq( + sq, + ), + ) + + return qb.getMany(ctx, q) +} + +// returns nil, sql.ErrNoRows if not found +func (qb *ImageStore) get(ctx context.Context, q *goqu.SelectDataset) (*models.Image, error) { + ret, err := qb.getMany(ctx, q) + if err != nil { + return nil, err + } + + if len(ret) == 0 { + return nil, sql.ErrNoRows + } + + return ret[0], nil +} + +func (qb *ImageStore) getMany(ctx context.Context, q *goqu.SelectDataset) ([]*models.Image, error) { + const single = false + var ret []*models.Image + var lastID int + if err := queryFunc(ctx, q, single, func(r *sqlx.Rows) error { + var f imageQueryRow + if err := r.StructScan(&f); err != nil { + return err + } + + i := f.resolve() + + if i.ID == lastID { + return fmt.Errorf("internal error: multiple rows returned for single image id %d", i.ID) + } + lastID = i.ID + + ret = append(ret, i) + return nil + }); err != nil { + return nil, err + } + + return ret, nil +} + +// Returns the custom cover for the gallery, if one has been set. +func (qb *ImageStore) CoverByGalleryID(ctx context.Context, galleryID int) (*models.Image, error) { + table := qb.table() + + sq := dialect.From(table). + InnerJoin( + galleriesImagesJoinTable, + goqu.On(table.Col(idColumn).Eq(galleriesImagesJoinTable.Col(imageIDColumn))), + ). + Select(table.Col(idColumn)). + Where(goqu.And( + galleriesImagesJoinTable.Col("gallery_id").Eq(galleryID), + galleriesImagesJoinTable.Col("cover").IsTrue(), + )) + + q := qb.selectDataset().Prepared(true).Where( + table.Col(idColumn).Eq( + sq, + ), + ) + + ret, err := qb.getMany(ctx, q) + if err != nil { + return nil, fmt.Errorf("getting cover for gallery %d: %w", galleryID, err) + } + + switch { + case len(ret) > 1: + return nil, fmt.Errorf("internal error: multiple covers returned for gallery %d", galleryID) + case len(ret) == 1: + return ret[0], nil + default: + return nil, nil + } +} + +func (qb *ImageStore) GetFiles(ctx context.Context, id int) ([]models.File, error) { + fileIDs, err := imageRepository.files.get(ctx, id) + if err != nil { + return nil, err + } + + // use fileStore to load files + files, err := qb.repo.File.Find(ctx, fileIDs...) + if err != nil { + return nil, err + } + + ret := make([]models.File, len(files)) + copy(ret, files) + + return ret, nil +} + +func (qb *ImageStore) GetManyFileIDs(ctx context.Context, ids []int) ([][]models.FileID, error) { + const primaryOnly = false + return imageRepository.files.getMany(ctx, ids, primaryOnly) +} + +func (qb *ImageStore) FindByFileID(ctx context.Context, fileID models.FileID) ([]*models.Image, error) { + table := qb.table() + + sq := dialect.From(table). + InnerJoin( + imagesFilesJoinTable, + goqu.On(table.Col(idColumn).Eq(imagesFilesJoinTable.Col(imageIDColumn))), + ). + Select(table.Col(idColumn)).Where(imagesFilesJoinTable.Col(fileIDColumn).Eq(fileID)) + + ret, err := qb.findBySubquery(ctx, sq) + if err != nil { + return nil, fmt.Errorf("getting image by file id %d: %w", fileID, err) + } + + return ret, nil +} + +func (qb *ImageStore) CountByFileID(ctx context.Context, fileID models.FileID) (int, error) { + joinTable := imagesFilesJoinTable + + q := dialect.Select(goqu.COUNT("*")).From(joinTable).Where(joinTable.Col(fileIDColumn).Eq(fileID)) + return count(ctx, q) +} + +func (qb *ImageStore) FindByFingerprints(ctx context.Context, fp []models.Fingerprint) ([]*models.Image, error) { + table := qb.table() + fingerprintTable := fingerprintTableMgr.table + + var ex []exp.Expression + + for _, v := range fp { + ex = append(ex, goqu.And( + fingerprintTable.Col("type").Eq(v.Type), + fingerprintTable.Col("fingerprint").Eq(v.Fingerprint), + )) + } + + sq := dialect.From(table). + InnerJoin( + imagesFilesJoinTable, + goqu.On(table.Col(idColumn).Eq(imagesFilesJoinTable.Col(imageIDColumn))), + ). + InnerJoin( + fingerprintTable, + goqu.On(fingerprintTable.Col(fileIDColumn).Eq(imagesFilesJoinTable.Col(fileIDColumn))), + ). + Select(table.Col(idColumn)).Where(goqu.Or(ex...)) + + ret, err := qb.findBySubquery(ctx, sq) + if err != nil { + return nil, fmt.Errorf("getting image by fingerprints: %w", err) + } + + return ret, nil +} + +func (qb *ImageStore) FindByChecksum(ctx context.Context, checksum string) ([]*models.Image, error) { + return qb.FindByFingerprints(ctx, []models.Fingerprint{ + { + Type: models.FingerprintTypeMD5, + Fingerprint: checksum, + }, + }) +} + +var defaultGalleryOrder = []exp.OrderedExpression{ + goqu.L("COALESCE(folders.path, '') || COALESCE(files.basename, '') COLLATE NATURAL_CI").Asc(), + goqu.L("COALESCE(images.title, cast(images.id as text)) COLLATE NATURAL_CI").Asc(), +} + +func (qb *ImageStore) FindByGalleryID(ctx context.Context, galleryID int) ([]*models.Image, error) { + table := qb.table() + + sq := dialect.From(table). + InnerJoin( + galleriesImagesJoinTable, + goqu.On(table.Col(idColumn).Eq(galleriesImagesJoinTable.Col(imageIDColumn))), + ). + Select(table.Col(idColumn)).Where( + galleriesImagesJoinTable.Col("gallery_id").Eq(galleryID), + ) + + q := qb.selectDataset().Prepared(true).Where( + table.Col(idColumn).Eq( + sq, + ), + ).Order(defaultGalleryOrder...) + + ret, err := qb.getMany(ctx, q) + if err != nil { + return nil, fmt.Errorf("getting images for gallery %d: %w", galleryID, err) + } + + return ret, nil +} + +func (qb *ImageStore) FindByGalleryIDIndex(ctx context.Context, galleryID int, index uint) (*models.Image, error) { + table := qb.table() + + q := qb.selectDataset(). + InnerJoin( + galleriesImagesJoinTable, + goqu.On(table.Col(idColumn).Eq(galleriesImagesJoinTable.Col(imageIDColumn))), + ). + Where(galleriesImagesJoinTable.Col(galleryIDColumn).Eq(galleryID)). + Prepared(true). + Order(defaultGalleryOrder...). + Limit(1).Offset(index) + + ret, err := qb.getMany(ctx, q) + if err != nil { + return nil, fmt.Errorf("getting images for gallery %d: %w", galleryID, err) + } + + if len(ret) == 0 { + return nil, nil + } + + return ret[0], nil +} + +func (qb *ImageStore) CountByGalleryID(ctx context.Context, galleryID int) (int, error) { + joinTable := goqu.T(galleriesImagesTable) + + q := dialect.Select(goqu.COUNT("*")).From(joinTable).Where(joinTable.Col("gallery_id").Eq(galleryID)) + return count(ctx, q) +} + +func (qb *ImageStore) OCountByPerformerID(ctx context.Context, performerID int) (int, error) { + table := qb.table() + joinTable := performersImagesJoinTable + q := dialect.Select(goqu.COALESCE(goqu.SUM("o_counter"), 0)).From(table).InnerJoin(joinTable, goqu.On(table.Col(idColumn).Eq(joinTable.Col(imageIDColumn)))).Where(joinTable.Col(performerIDColumn).Eq(performerID)) + + var ret int + if err := querySimple(ctx, q, &ret); err != nil { + return 0, err + } + + return ret, nil +} + +func (qb *ImageStore) OCountByStudioID(ctx context.Context, studioID int) (int, error) { + table := qb.table() + q := dialect.Select(goqu.COALESCE(goqu.SUM("o_counter"), 0)).From(table).Where( + table.Col(studioIDColumn).Eq(studioID), + ) + + var ret int + if err := querySimple(ctx, q, &ret); err != nil { + return 0, err + } + + return ret, nil +} + +func (qb *ImageStore) OCount(ctx context.Context) (int, error) { + table := qb.table() + + q := dialect.Select(goqu.COALESCE(goqu.SUM("o_counter"), 0)).From(table) + var ret int + if err := querySimple(ctx, q, &ret); err != nil { + return 0, err + } + + return ret, nil +} + +func (qb *ImageStore) FindByFolderID(ctx context.Context, folderID models.FolderID) ([]*models.Image, error) { + table := qb.table() + fileTable := goqu.T(fileTable) + + sq := dialect.From(table). + InnerJoin( + imagesFilesJoinTable, + goqu.On(table.Col(idColumn).Eq(imagesFilesJoinTable.Col(imageIDColumn))), + ). + InnerJoin( + fileTable, + goqu.On(imagesFilesJoinTable.Col(fileIDColumn).Eq(fileTable.Col(idColumn))), + ). + Select(table.Col(idColumn)).Where( + fileTable.Col("parent_folder_id").Eq(folderID), + ) + + ret, err := qb.findBySubquery(ctx, sq) + if err != nil { + return nil, fmt.Errorf("getting image by folder: %w", err) + } + + return ret, nil +} + +func (qb *ImageStore) FindByZipFileID(ctx context.Context, zipFileID models.FileID) ([]*models.Image, error) { + table := qb.table() + fileTable := goqu.T(fileTable) + + sq := dialect.From(table). + InnerJoin( + imagesFilesJoinTable, + goqu.On(table.Col(idColumn).Eq(imagesFilesJoinTable.Col(imageIDColumn))), + ). + InnerJoin( + fileTable, + goqu.On(imagesFilesJoinTable.Col(fileIDColumn).Eq(fileTable.Col(idColumn))), + ). + Select(table.Col(idColumn)).Where( + fileTable.Col("zip_file_id").Eq(zipFileID), + ) + + ret, err := qb.findBySubquery(ctx, sq) + if err != nil { + return nil, fmt.Errorf("getting image by zip file: %w", err) + } + + return ret, nil +} + +func (qb *ImageStore) Count(ctx context.Context) (int, error) { + q := dialect.Select(goqu.COUNT("*")).From(qb.table()) + return count(ctx, q) +} + +func (qb *ImageStore) Size(ctx context.Context) (float64, error) { + table := qb.table() + fileTable := fileTableMgr.table + q := dialect.Select( + goqu.COALESCE(goqu.SUM(fileTableMgr.table.Col("size")), 0), + ).From(table).InnerJoin( + imagesFilesJoinTable, + goqu.On(table.Col(idColumn).Eq(imagesFilesJoinTable.Col(imageIDColumn))), + ).InnerJoin( + fileTable, + goqu.On(imagesFilesJoinTable.Col(fileIDColumn).Eq(fileTable.Col(idColumn))), + ) + var ret float64 + if err := querySimple(ctx, q, &ret); err != nil { + return 0, err + } + + return ret, nil +} + +func (qb *ImageStore) All(ctx context.Context) ([]*models.Image, error) { + return qb.getMany(ctx, qb.selectDataset()) +} + +func (qb *ImageStore) makeQuery(ctx context.Context, imageFilter *models.ImageFilterType, findFilter *models.FindFilterType) (*queryBuilder, error) { + if imageFilter == nil { + imageFilter = &models.ImageFilterType{} + } + if findFilter == nil { + findFilter = &models.FindFilterType{} + } + + query := imageRepository.newQuery() + distinctIDs(&query, imageTable) + + if q := findFilter.Q; q != nil && *q != "" { + query.addJoins( + join{ + table: imagesFilesTable, + onClause: "images_files.image_id = images.id AND images_files.\"primary\" = true", + }, + join{ + table: fileTable, + onClause: "images_files.file_id = files.id", + }, + join{ + table: folderTable, + onClause: "files.parent_folder_id = folders.id", + }, + join{ + table: fingerprintTable, + onClause: "files_fingerprints.file_id = images_files.file_id", + }, + ) + + filepathColumn := "folders.path || '" + string(filepath.Separator) + "' || files.basename" + searchColumns := []string{"images.title", filepathColumn, "files_fingerprints.fingerprint"} + query.parseQueryString(searchColumns, *q) + } + + filter := filterBuilderFromHandler(ctx, &imageFilterHandler{ + imageFilter: imageFilter, + }) + + if err := query.addFilter(filter); err != nil { + return nil, err + } + + if err := qb.setImageSortAndPagination(&query, findFilter); err != nil { + return nil, err + } + + return &query, nil +} + +func (qb *ImageStore) Query(ctx context.Context, options models.ImageQueryOptions) (*models.ImageQueryResult, error) { + query, err := qb.makeQuery(ctx, options.ImageFilter, options.FindFilter) + if err != nil { + return nil, err + } + + result, err := qb.queryGroupedFields(ctx, options, *query) + if err != nil { + return nil, fmt.Errorf("error querying aggregate fields: %w", err) + } + + return result, nil +} + +func (qb *ImageStore) queryGroupedFields(ctx context.Context, options models.ImageQueryOptions, query queryBuilder) (*models.ImageQueryResult, error) { + if options.Count { + query.addColumn("COUNT(*) OVER () AS total_count") + } + if options.Megapixels { + query.addJoins( + join{ + table: imagesFilesTable, + onClause: "images_files.image_id = images.id AND images_files.\"primary\" = true", + }, + join{ + table: imageFileTable, + onClause: "images_files.file_id = image_files.file_id", + }, + ) + query.addColumn(` +COALESCE( + SUM( + COALESCE(image_files.width, 0) * COALESCE(image_files.height, 0) + ) OVER (), 0 +) / 1000000 AS total_custom +`) + query.addGroupBy("image_files.width", "image_files.height") + } + if options.TotalSize { + query.addJoins( + join{ + table: imagesFilesTable, + onClause: "images_files.image_id = images.id AND images_files.\"primary\" = true", + }, + join{ + table: fileTable, + onClause: "images_files.file_id = files.id", + }, + ) + query.addColumn("SUM(COALESCE(files.size, 0)) OVER () AS total_size") + query.addGroupBy("files.size") + } + + // Support counting only + const includeSortPagination = true + var countOnly = options.FindFilter != nil && options.FindFilter.IsCounting() + + // Execute aggregate query + var obj *RowsWithCounts + var err error + if obj, err = sceneRepository.runIdsWithCount(ctx, query.toSQL(includeSortPagination), query.args); err != nil { + return nil, err + } + + ret := models.NewImageQueryResult(qb) + if len(obj.IDs) == 0 { + return ret, nil + } + + if !countOnly { + ret.IDs = obj.IDs + } + if options.Count { + ret.Count = int(obj.TotalCount.Int64) + } + if options.TotalSize { + ret.TotalSize = obj.TotalSize.Float64 + } + if options.Megapixels { + ret.Megapixels = obj.TotalCustom.Float64 + } + return ret, nil +} + +func (qb *ImageStore) QueryCount(ctx context.Context, imageFilter *models.ImageFilterType, findFilter *models.FindFilterType) (int, error) { + query, err := qb.makeQuery(ctx, imageFilter, findFilter) + if err != nil { + return 0, err + } + + return query.executeCount(ctx) +} + +var imageSortOptions = sortOptions{ + "created_at", + "date", + "file_count", + "file_mod_time", + "filesize", + "id", + "o_counter", + "path", + "performer_count", + "random", + "rating", + "tag_count", + "title", + "updated_at", +} + +func (qb *ImageStore) setImageSortAndPagination(q *queryBuilder, findFilter *models.FindFilterType) error { + sortClause := "" + + if findFilter != nil && findFilter.Sort != nil && *findFilter.Sort != "" { + sort := findFilter.GetSort("title") + direction := findFilter.GetDirection() + + // CVE-2024-32231 - ensure sort is in the list of allowed sorts + if err := imageSortOptions.validateSort(sort); err != nil { + return err + } + + // translate sort field + if sort == "file_mod_time" { + sort = "mod_time" + } + + addFilesJoin := func() { + q.addJoins( + join{ + sort: true, + table: imagesFilesTable, + onClause: "images_files.image_id = images.id AND images_files.\"primary\" = true", + }, + join{ + sort: true, + table: fileTable, + onClause: "images_files.file_id = files.id", + }, + ) + } + + addFolderJoin := func() { + q.addJoins(join{ + sort: true, + table: folderTable, + onClause: "files.parent_folder_id = folders.id", + }) + } + + switch sort { + case "path": + addFilesJoin() + addFolderJoin() + sortClause = " ORDER BY COALESCE(folders.path, '') || COALESCE(files.basename, '') COLLATE NATURAL_CI " + direction + q.addGroupBy("folders.path", "files.basename") + case "file_count": + sortClause = getCountSort(imageTable, imagesFilesTable, imageIDColumn, direction) + case "tag_count": + sortClause = getCountSort(imageTable, imagesTagsTable, imageIDColumn, direction) + case "performer_count": + sortClause = getCountSort(imageTable, performersImagesTable, imageIDColumn, direction) + case "mod_time", "filesize": + addFilesJoin() + add, agg := getSort(sort, direction, "files") + sortClause = add + q.addGroupBy(agg...) + case "resolution": + addFilesJoin() + q.addJoins(join{ + sort: true, + table: imageFileTable, + onClause: "images_files.file_id = image_files.file_id", + }) + sortClause = " ORDER BY MIN(image_files.width, image_files.height) " + direction + case "title": + addFilesJoin() + addFolderJoin() + sortClause = " ORDER BY COALESCE(images.title, files.basename) COLLATE NATURAL_CI " + direction + ", folders.path COLLATE NATURAL_CI " + direction + q.addGroupBy("images.title", "files.basename", "folders.path") + default: + add, agg := getSort(sort, direction, "images") + sortClause = add + q.addGroupBy(agg...) + } + + // Whatever the sorting, always use title/id as a final sort + sortClause += ", COALESCE(images.title, CAST(images.id as text)) COLLATE NATURAL_CI ASC" + q.addGroupBy("images.title", "images.id") + } + + q.sort = sortClause + q.pagination = getPagination(findFilter) + + return nil +} + +func (qb *ImageStore) AddFileID(ctx context.Context, id int, fileID models.FileID) error { + const firstPrimary = false + return imagesFilesTableMgr.insertJoins(ctx, id, firstPrimary, []models.FileID{fileID}) +} + +// RemoveFileID removes the file ID from the image. +// If the file ID is the primary file, then the next file in the list is set as the primary file. +func (qb *ImageStore) RemoveFileID(ctx context.Context, id int, fileID models.FileID) error { + fileIDs, err := imagesFilesTableMgr.get(ctx, id) + if err != nil { + return fmt.Errorf("getting file IDs for image %d: %w", id, err) + } + + fileIDs = sliceutil.Filter(fileIDs, func(f models.FileID) bool { + return f != fileID + }) + + return imagesFilesTableMgr.replaceJoins(ctx, id, fileIDs) +} + +func (qb *ImageStore) GetGalleryIDs(ctx context.Context, imageID int) ([]int, error) { + return imageRepository.galleries.getIDs(ctx, imageID) +} + +// func (qb *imageQueryBuilder) UpdateGalleries(ctx context.Context, imageID int, galleryIDs []int) error { +// // Delete the existing joins and then create new ones +// return qb.galleriesRepository().replace(ctx, imageID, galleryIDs) +// } + +func (qb *ImageStore) GetPerformerIDs(ctx context.Context, imageID int) ([]int, error) { + return imageRepository.performers.getIDs(ctx, imageID) +} + +func (qb *ImageStore) UpdatePerformers(ctx context.Context, imageID int, performerIDs []int) error { + // Delete the existing joins and then create new ones + return imageRepository.performers.replace(ctx, imageID, performerIDs) +} + +func (qb *ImageStore) GetTagIDs(ctx context.Context, imageID int) ([]int, error) { + return imageRepository.tags.getIDs(ctx, imageID) +} + +func (qb *ImageStore) UpdateTags(ctx context.Context, imageID int, tagIDs []int) error { + // Delete the existing joins and then create new ones + return imageRepository.tags.replace(ctx, imageID, tagIDs) +} + +func (qb *ImageStore) GetURLs(ctx context.Context, imageID int) ([]string, error) { + return imagesURLsTableMgr.get(ctx, imageID) +} diff --git a/pkg/postgres/image_filter.go b/pkg/postgres/image_filter.go new file mode 100644 index 0000000000..3aca2f5db7 --- /dev/null +++ b/pkg/postgres/image_filter.go @@ -0,0 +1,307 @@ +package postgres + +import ( + "context" + + "github.com/stashapp/stash/pkg/models" +) + +type imageFilterHandler struct { + imageFilter *models.ImageFilterType +} + +func (qb *imageFilterHandler) validate() error { + imageFilter := qb.imageFilter + if imageFilter == nil { + return nil + } + + if err := validateFilterCombination(imageFilter.OperatorFilter); err != nil { + return err + } + + if subFilter := imageFilter.SubFilter(); subFilter != nil { + sqb := &imageFilterHandler{imageFilter: subFilter} + if err := sqb.validate(); err != nil { + return err + } + } + + return nil +} + +func (qb *imageFilterHandler) handle(ctx context.Context, f *filterBuilder) { + imageFilter := qb.imageFilter + if imageFilter == nil { + return + } + + if err := qb.validate(); err != nil { + f.setError(err) + return + } + + sf := imageFilter.SubFilter() + if sf != nil { + sub := &imageFilterHandler{sf} + handleSubFilter(ctx, sub, f, imageFilter.OperatorFilter) + } + + f.handleCriterion(ctx, qb.criterionHandler()) +} + +func (qb *imageFilterHandler) criterionHandler() criterionHandler { + imageFilter := qb.imageFilter + return compoundHandler{ + intCriterionHandler(imageFilter.ID, "images.id", nil), + criterionHandlerFunc(func(ctx context.Context, f *filterBuilder) { + if imageFilter.Checksum != nil { + imageRepository.addImagesFilesTable(f) + f.addInnerJoin(fingerprintTable, "fingerprints_md5", "images_files.file_id = fingerprints_md5.file_id AND fingerprints_md5.type = 'md5'") + } + + stringCriterionHandler(imageFilter.Checksum, "fingerprints_md5.fingerprint")(ctx, f) + }), + stringCriterionHandler(imageFilter.Title, "images.title"), + stringCriterionHandler(imageFilter.Code, "images.code"), + stringCriterionHandler(imageFilter.Details, "images.details"), + stringCriterionHandler(imageFilter.Photographer, "images.photographer"), + + pathCriterionHandler(imageFilter.Path, "folders.path", "files.basename", imageRepository.addFoldersTable), + qb.fileCountCriterionHandler(imageFilter.FileCount), + intCriterionHandler(imageFilter.Rating100, "images.rating", nil), + intCriterionHandler(imageFilter.OCounter, "images.o_counter", nil), + boolCriterionHandler(imageFilter.Organized, "images.organized", nil), + &dateCriterionHandler{imageFilter.Date, "images.date", nil}, + qb.urlsCriterionHandler(imageFilter.URL), + + resolutionCriterionHandler(imageFilter.Resolution, "image_files.height", "image_files.width", imageRepository.addImageFilesTable), + orientationCriterionHandler(imageFilter.Orientation, "image_files.height", "image_files.width", imageRepository.addImageFilesTable), + qb.missingCriterionHandler(imageFilter.IsMissing), + + qb.tagsCriterionHandler(imageFilter.Tags), + qb.tagCountCriterionHandler(imageFilter.TagCount), + qb.galleriesCriterionHandler(imageFilter.Galleries), + qb.performersCriterionHandler(imageFilter.Performers), + qb.performerCountCriterionHandler(imageFilter.PerformerCount), + studioCriterionHandler(imageTable, imageFilter.Studios), + qb.performerTagsCriterionHandler(imageFilter.PerformerTags), + qb.performerFavoriteCriterionHandler(imageFilter.PerformerFavorite), + qb.performerAgeCriterionHandler(imageFilter.PerformerAge), + ×tampCriterionHandler{imageFilter.CreatedAt, "images.created_at", nil}, + ×tampCriterionHandler{imageFilter.UpdatedAt, "images.updated_at", nil}, + + &relatedFilterHandler{ + relatedIDCol: "galleries_images.gallery_id", + relatedRepo: galleryRepository.repository, + relatedHandler: &galleryFilterHandler{imageFilter.GalleriesFilter}, + joinFn: func(f *filterBuilder) { + imageRepository.galleries.innerJoin(f, "", "images.id") + }, + }, + + &relatedFilterHandler{ + relatedIDCol: "performers_join.performer_id", + relatedRepo: performerRepository.repository, + relatedHandler: &performerFilterHandler{imageFilter.PerformersFilter}, + joinFn: func(f *filterBuilder) { + imageRepository.performers.innerJoin(f, "performers_join", "images.id") + }, + }, + + &relatedFilterHandler{ + relatedIDCol: "images.studio_id", + relatedRepo: studioRepository.repository, + relatedHandler: &studioFilterHandler{imageFilter.StudiosFilter}, + }, + + &relatedFilterHandler{ + relatedIDCol: "image_tag.tag_id", + relatedRepo: tagRepository.repository, + relatedHandler: &tagFilterHandler{imageFilter.TagsFilter}, + joinFn: func(f *filterBuilder) { + imageRepository.tags.innerJoin(f, "image_tag", "images.id") + }, + }, + + &relatedFilterHandler{ + relatedIDCol: "files.id", + relatedRepo: fileRepository.repository, + relatedHandler: &fileFilterHandler{ + fileFilter: imageFilter.FilesFilter, + isRelated: true, + }, + joinFn: func(f *filterBuilder) { + imageRepository.addFilesTable(f) + imageRepository.addFoldersTable(f) + }, + // don't use a subquery; join directly + directJoin: true, + }, + } +} + +func (qb *imageFilterHandler) fileCountCriterionHandler(fileCount *models.IntCriterionInput) criterionHandlerFunc { + h := countCriterionHandlerBuilder{ + primaryTable: imageTable, + joinTable: imagesFilesTable, + primaryFK: imageIDColumn, + } + + return h.handler(fileCount) +} + +func (qb *imageFilterHandler) missingCriterionHandler(isMissing *string) criterionHandlerFunc { + return func(ctx context.Context, f *filterBuilder) { + if isMissing != nil && *isMissing != "" { + switch *isMissing { + case "studio": + f.addWhere("images.studio_id IS NULL") + case "performers": + imageRepository.performers.join(f, "performers_join", "images.id") + f.addWhere("performers_join.image_id IS NULL") + case "galleries": + imageRepository.galleries.join(f, "galleries_join", "images.id") + f.addWhere("galleries_join.image_id IS NULL") + case "tags": + imageRepository.tags.join(f, "tags_join", "images.id") + f.addWhere("tags_join.image_id IS NULL") + default: + f.addWhere("(images." + *isMissing + " IS NULL OR TRIM(CAST(images." + *isMissing + " AS TEXT)) = '')") + } + } + } +} + +func (qb *imageFilterHandler) urlsCriterionHandler(url *models.StringCriterionInput) criterionHandlerFunc { + h := stringListCriterionHandlerBuilder{ + primaryTable: imageTable, + primaryFK: imageIDColumn, + joinTable: imagesURLsTable, + stringColumn: imageURLColumn, + addJoinTable: func(f *filterBuilder) { + imagesURLsTableMgr.join(f, "", "images.id") + }, + } + + return h.handler(url) +} + +func (qb *imageFilterHandler) getMultiCriterionHandlerBuilder(foreignTable, joinTable, foreignFK string, addJoinsFunc func(f *filterBuilder)) multiCriterionHandlerBuilder { + return multiCriterionHandlerBuilder{ + primaryTable: imageTable, + foreignTable: foreignTable, + joinTable: joinTable, + primaryFK: imageIDColumn, + foreignFK: foreignFK, + addJoinsFunc: addJoinsFunc, + } +} + +func (qb *imageFilterHandler) tagsCriterionHandler(tags *models.HierarchicalMultiCriterionInput) criterionHandlerFunc { + h := joinedHierarchicalMultiCriterionHandlerBuilder{ + primaryTable: imageTable, + foreignTable: tagTable, + foreignFK: "tag_id", + + relationsTable: "tags_relations", + joinAs: "image_tag", + joinTable: imagesTagsTable, + primaryFK: imageIDColumn, + } + + return h.handler(tags) +} + +func (qb *imageFilterHandler) tagCountCriterionHandler(tagCount *models.IntCriterionInput) criterionHandlerFunc { + h := countCriterionHandlerBuilder{ + primaryTable: imageTable, + joinTable: imagesTagsTable, + primaryFK: imageIDColumn, + } + + return h.handler(tagCount) +} + +func (qb *imageFilterHandler) galleriesCriterionHandler(galleries *models.MultiCriterionInput) criterionHandlerFunc { + addJoinsFunc := func(f *filterBuilder) { + if galleries.Modifier == models.CriterionModifierIncludes || galleries.Modifier == models.CriterionModifierIncludesAll { + f.addInnerJoin(galleriesImagesTable, "", "galleries_images.image_id = images.id") + f.addInnerJoin(galleryTable, "", "galleries_images.gallery_id = galleries.id") + } + } + h := qb.getMultiCriterionHandlerBuilder(galleryTable, galleriesImagesTable, galleryIDColumn, addJoinsFunc) + + return h.handler(galleries) +} + +func (qb *imageFilterHandler) performersCriterionHandler(performers *models.MultiCriterionInput) criterionHandlerFunc { + h := joinedMultiCriterionHandlerBuilder{ + primaryTable: imageTable, + joinTable: performersImagesTable, + joinAs: "performers_join", + primaryFK: imageIDColumn, + foreignFK: performerIDColumn, + + addJoinTable: func(f *filterBuilder) { + imageRepository.performers.join(f, "performers_join", "images.id") + }, + } + + return h.handler(performers) +} + +func (qb *imageFilterHandler) performerCountCriterionHandler(performerCount *models.IntCriterionInput) criterionHandlerFunc { + h := countCriterionHandlerBuilder{ + primaryTable: imageTable, + joinTable: performersImagesTable, + primaryFK: imageIDColumn, + } + + return h.handler(performerCount) +} + +func (qb *imageFilterHandler) performerFavoriteCriterionHandler(performerfavorite *bool) criterionHandlerFunc { + return func(ctx context.Context, f *filterBuilder) { + if performerfavorite != nil { + f.addLeftJoin("performers_images", "", "images.id = performers_images.image_id") + + if *performerfavorite { + // contains at least one favorite + f.addLeftJoin("performers", "", "performers.id = performers_images.performer_id") + f.addWhere("performers.favorite = true") + } else { + // contains zero favorites + f.addLeftJoin(`(SELECT performers_images.image_id as id FROM performers_images +JOIN performers ON performers.id = performers_images.performer_id +GROUP BY performers_images.image_id HAVING SUM(performers.favorite) = false)`, "nofaves", "images.id = nofaves.id") + f.addWhere("performers_images.image_id IS NULL OR nofaves.id IS NOT NULL") + } + } + } +} + +func (qb *imageFilterHandler) performerAgeCriterionHandler(performerAge *models.IntCriterionInput) criterionHandlerFunc { + return func(ctx context.Context, f *filterBuilder) { + if performerAge != nil { + f.addInnerJoin("performers_images", "", "images.id = performers_images.image_id") + f.addInnerJoin("performers", "", "performers_images.performer_id = performers.id") + + f.addWhere("images.date != '' AND performers.birthdate != ''") + f.addWhere("images.date IS NOT NULL AND performers.birthdate IS NOT NULL") + + ageCalc := "EXTRACT(YEAR FROM AGE(images.date, performers.birthdate))" + whereClause, args := getIntWhereClause(ageCalc, performerAge.Modifier, performerAge.Value, performerAge.Value2) + f.addWhere(whereClause, args...) + } + } +} + +func (qb *imageFilterHandler) performerTagsCriterionHandler(tags *models.HierarchicalMultiCriterionInput) criterionHandler { + return &joinedPerformerTagsHandler{ + criterion: tags, + primaryTable: imageTable, + joinTable: performersImagesTable, + joinPrimaryKey: imageIDColumn, + } +} diff --git a/pkg/postgres/interfaces.go b/pkg/postgres/interfaces.go new file mode 100644 index 0000000000..4992c8e106 --- /dev/null +++ b/pkg/postgres/interfaces.go @@ -0,0 +1,46 @@ +package postgres + +import "github.com/stashapp/stash/pkg/database" + +func (db *Database) Blobs() database.BlobStore { + return db.storeRepository.Blobs +} +func (db *Database) File() database.FileStore { + return db.storeRepository.File +} +func (db *Database) Folder() database.FolderStore { + return db.storeRepository.Folder +} +func (db *Database) Image() database.ImageStore { + return db.storeRepository.Image +} +func (db *Database) Gallery() database.GalleryStore { + return db.storeRepository.Gallery +} +func (db *Database) GalleryChapter() database.GalleryChapterStore { + return db.storeRepository.GalleryChapter +} +func (db *Database) Scene() database.SceneStore { + return db.storeRepository.Scene +} +func (db *Database) SceneMarker() database.SceneMarkerStore { + return db.storeRepository.SceneMarker +} +func (db *Database) Performer() database.PerformerStore { + return db.storeRepository.Performer +} +func (db *Database) SavedFilter() database.SavedFilterStore { + return db.storeRepository.SavedFilter +} +func (db *Database) Studio() database.StudioStore { + return db.storeRepository.Studio +} +func (db *Database) Tag() database.TagStore { + return db.storeRepository.Tag +} +func (db *Database) Group() database.GroupStore { + return db.storeRepository.Group +} +func (db *Database) NewMigrator() (database.MigrateStore, error) { + return NewMigrator(db) +} diff --git a/pkg/postgres/migrate.go b/pkg/postgres/migrate.go new file mode 100644 index 0000000000..38f0b95b13 --- /dev/null +++ b/pkg/postgres/migrate.go @@ -0,0 +1,190 @@ +package postgres + +import ( + "context" + "fmt" + + "github.com/golang-migrate/migrate/v4" + postgresmig "github.com/golang-migrate/migrate/v4/database/postgres" + "github.com/golang-migrate/migrate/v4/source/iofs" + "github.com/jmoiron/sqlx" + "github.com/stashapp/stash/pkg/logger" +) + +func (db *Database) needsMigration() bool { + return db.schemaVersion != appSchemaVersion +} + +type Migrator struct { + db *Database + conn *sqlx.DB + m *migrate.Migrate +} + +func NewMigrator(db *Database) (*Migrator, error) { + m := &Migrator{ + db: db, + } + + const disableForeignKeys = true + const writable = true + var err error + m.conn, err = m.db.open(disableForeignKeys, writable) + if err != nil { + return nil, err + } + + m.conn.SetMaxOpenConns(maxReadConnections) + m.conn.SetMaxIdleConns(maxReadConnections) + m.conn.SetConnMaxIdleTime(dbConnTimeout) + + m.m, err = m.getMigrate() + + // if error encountered, close the connection + if err != nil { + m.Close() + } + + return m, err +} + +func (m *Migrator) Close() { + if m.m != nil { + m.m.Close() + m.m = nil + } +} + +func (m *Migrator) CurrentSchemaVersion() uint { + databaseSchemaVersion, _, _ := m.m.Version() + return databaseSchemaVersion +} + +func (m *Migrator) RequiredSchemaVersion() uint { + return appSchemaVersion +} + +func (m *Migrator) getMigrate() (*migrate.Migrate, error) { + migrations, err := iofs.New(migrationsBox, "migrations") + if err != nil { + return nil, err + } + + driver, err := postgresmig.WithInstance(m.conn.DB, &postgresmig.Config{}) + if err != nil { + return nil, err + } + + // use sqlite3Driver so that migration has access to durationToTinyInt + return migrate.NewWithInstance( + "iofs", + migrations, + "postgres", + driver, + ) +} + +func (m *Migrator) RunMigration(ctx context.Context, newVersion uint) error { + databaseSchemaVersion, _, _ := m.m.Version() + + if newVersion != databaseSchemaVersion+1 { + return fmt.Errorf("invalid migration version %d, expected %d", newVersion, databaseSchemaVersion+1) + } + + // run pre migrations as needed + if err := m.runCustomMigrations(ctx, preMigrations[newVersion]); err != nil { + return fmt.Errorf("running pre migrations for schema version %d: %w", newVersion, err) + } + + if err := m.m.Steps(1); err != nil { + // migration failed + return err + } + + // run post migrations as needed + if err := m.runCustomMigrations(ctx, postMigrations[newVersion]); err != nil { + return fmt.Errorf("running post migrations for schema version %d: %w", newVersion, err) + } + + // update the schema version + m.db.schemaVersion, _, _ = m.m.Version() + + return nil +} + +func (m *Migrator) runCustomMigrations(ctx context.Context, fns []customMigrationFunc) error { + for _, fn := range fns { + if err := m.runCustomMigration(ctx, fn); err != nil { + return err + } + } + + return nil +} + +func (m *Migrator) runCustomMigration(ctx context.Context, fn customMigrationFunc) error { + if err := fn(ctx, m.conn); err != nil { + return err + } + + return nil +} + +func (m *Migrator) PostMigrate(ctx context.Context) error { + // optimise the database + var err error + logger.Info("Running database analyze") + + // don't use Optimize/vacuum as this adds a significant amount of time + // to the migration + err = analyze(ctx, m.conn) + + if err != nil { + return fmt.Errorf("error optimising database: %s", err) + } + + return nil +} + +func (db *Database) getDatabaseSchemaVersion() (uint, error) { + m, err := NewMigrator(db) + if err != nil { + return 0, err + } + defer m.Close() + + ret, _, _ := m.m.Version() + return ret, nil +} + +func (db *Database) ReInitialise() error { + return db.initialise() +} + +// RunAllMigrations runs all migrations to bring the database up to the current schema version +func (db *Database) RunAllMigrations() error { + ctx := context.Background() + + m, err := NewMigrator(db) + if err != nil { + return err + } + defer m.Close() + + databaseSchemaVersion, _, _ := m.m.Version() + stepNumber := appSchemaVersion - databaseSchemaVersion + if stepNumber != 0 { + logger.Infof("Migrating database from version %d to %d", databaseSchemaVersion, appSchemaVersion) + + // run each migration individually, and run custom migrations as needed + var i uint = 1 + for ; i <= stepNumber; i++ { + newVersion := databaseSchemaVersion + i + if err := m.RunMigration(ctx, newVersion); err != nil { + return err + } + } + } + + return nil +} diff --git a/pkg/postgres/migrations/10_sql_phash.up.sql b/pkg/postgres/migrations/10_sql_phash.up.sql new file mode 100644 index 0000000000..1ffa4e9d40 --- /dev/null +++ b/pkg/postgres/migrations/10_sql_phash.up.sql @@ -0,0 +1,7 @@ +CREATE OR REPLACE FUNCTION phash_distance(lhash bigint, rhash bigint) +RETURNS bigint +LANGUAGE sql +IMMUTABLE +AS $$ + SELECT length(replace(((lhash::bit(64) # rhash::bit(64))::text), '0', '')); +$$; diff --git a/pkg/postgres/migrations/11_studio_urls.up.sql b/pkg/postgres/migrations/11_studio_urls.up.sql new file mode 100644 index 0000000000..0fa21d5d7b --- /dev/null +++ b/pkg/postgres/migrations/11_studio_urls.up.sql @@ -0,0 +1,25 @@ +-- 73_studio_urls.up.sql +CREATE TABLE "studio_urls" ( + "studio_id" integer NOT NULL, + "position" integer NOT NULL, + "url" text NOT NULL, + foreign key("studio_id") references "studios"("id") on delete CASCADE, + PRIMARY KEY("studio_id", "position", "url") +); + +CREATE INDEX "studio_urls_url" on "studio_urls" ("url"); + +INSERT INTO "studio_urls" + ( + "studio_id", + "position", + "url" + ) + SELECT + "id", + '0', + "url" + FROM "studios" + WHERE "studios"."url" IS NOT NULL AND "studios"."url" != ''; + +ALTER TABLE "studios" DROP COLUMN "url"; diff --git a/pkg/postgres/migrations/12_tag_stash_ids.up.sql b/pkg/postgres/migrations/12_tag_stash_ids.up.sql new file mode 100644 index 0000000000..cf37ea243c --- /dev/null +++ b/pkg/postgres/migrations/12_tag_stash_ids.up.sql @@ -0,0 +1,10 @@ +-- 74_tag_stash_ids.up.sql +CREATE TABLE "tag_stash_ids" ( + "tag_id" integer, + "endpoint" text, + "stash_id" uuid, + "updated_at" timestamp not null default '1970-01-01T00:00:00Z', + foreign key("tag_id") references "tags"("id") on delete CASCADE +); + +CREATE UNIQUE INDEX tag_stash_ids_unique_idx ON tag_stash_ids (tag_id, endpoint, stash_id); diff --git a/pkg/postgres/migrations/13_date_precision.up.sql b/pkg/postgres/migrations/13_date_precision.up.sql new file mode 100644 index 0000000000..e026b54d2d --- /dev/null +++ b/pkg/postgres/migrations/13_date_precision.up.sql @@ -0,0 +1,14 @@ +--75_date_precision.up.sql +ALTER TABLE "scenes" ADD COLUMN "date_precision" SMALLINT; +ALTER TABLE "images" ADD COLUMN "date_precision" SMALLINT; +ALTER TABLE "galleries" ADD COLUMN "date_precision" SMALLINT; +ALTER TABLE "groups" ADD COLUMN "date_precision" SMALLINT; +ALTER TABLE "performers" ADD COLUMN "birthdate_precision" SMALLINT; +ALTER TABLE "performers" ADD COLUMN "death_date_precision" SMALLINT; + +UPDATE "scenes" SET "date_precision" = 0 WHERE "date" IS NOT NULL; +UPDATE "images" SET "date_precision" = 0 WHERE "date" IS NOT NULL; +UPDATE "galleries" SET "date_precision" = 0 WHERE "date" IS NOT NULL; +UPDATE "groups" SET "date_precision" = 0 WHERE "date" IS NOT NULL; +UPDATE "performers" SET "birthdate_precision" = 0 WHERE "birthdate" IS NOT NULL; +UPDATE "performers" SET "death_date_precision" = 0 WHERE "death_date" IS NOT NULL; diff --git a/pkg/postgres/migrations/1_initial.up.sql b/pkg/postgres/migrations/1_initial.up.sql new file mode 100644 index 0000000000..659527a777 --- /dev/null +++ b/pkg/postgres/migrations/1_initial.up.sql @@ -0,0 +1,481 @@ +CREATE COLLATION IF NOT EXISTS NATURAL_CI (provider = icu, locale = 'en@colNumeric=yes'); +CREATE COLLATION IF NOT EXISTS NOCASE (provider = icu, locale = 'und-u-ks-level2', deterministic = false); +CREATE TABLE blobs ( + checksum varchar(255) NOT NULL PRIMARY KEY, + blob bytea +); +CREATE TABLE tags ( + id serial not null primary key, + name text, + created_at TIMESTAMP WITH TIME ZONE not null, + updated_at TIMESTAMP WITH TIME ZONE not null, + ignore_auto_tag boolean not null default FALSE, + description text, + image_blob varchar(255) REFERENCES blobs(checksum), + favorite boolean not null default false +); +CREATE TABLE folders ( + id serial not null primary key, + path text NOT NULL, + parent_folder_id integer, + mod_time TIMESTAMP WITH TIME ZONE not null, + created_at TIMESTAMP WITH TIME ZONE not null, + updated_at TIMESTAMP WITH TIME ZONE not null, + foreign key(parent_folder_id) references folders(id) on delete SET NULL +); +CREATE TABLE files ( + id serial not null primary key, + basename varchar(255) NOT NULL, + zip_file_id integer, + parent_folder_id integer not null, + size bigint NOT NULL, + mod_time TIMESTAMP WITH TIME ZONE not null, + created_at TIMESTAMP WITH TIME ZONE not null, + updated_at TIMESTAMP WITH TIME ZONE not null, + foreign key(zip_file_id) references files(id), + foreign key(parent_folder_id) references folders(id), + CHECK (basename != '') +); +ALTER TABLE folders ADD COLUMN zip_file_id integer REFERENCES files(id); +CREATE TABLE IF NOT EXISTS performers ( + id serial not null primary key, + name text not null, + disambiguation text, + gender varchar(20), + birthdate date, + ethnicity text, + country text, + eye_color text, + height int, + measurements text, + fake_tits text, + career_length text, + tattoos text, + piercings text, + favorite boolean not null default FALSE, + created_at TIMESTAMP WITH TIME ZONE not null, + updated_at TIMESTAMP WITH TIME ZONE not null, + details text, + death_date date, + hair_color text, + weight integer, + rating smallint, + ignore_auto_tag boolean not null default FALSE, + image_blob varchar(255) REFERENCES blobs(checksum), + penis_length float, + circumcised text +); +CREATE TABLE IF NOT EXISTS studios ( + id serial not null primary key, + name text NOT NULL, + url VARCHAR(2048), + parent_id INTEGER DEFAULT NULL REFERENCES studios(id) ON DELETE SET NULL, + created_at TIMESTAMP WITH TIME ZONE NOT NULL, + updated_at TIMESTAMP WITH TIME ZONE NOT NULL, + details TEXT, + rating smallint, + ignore_auto_tag BOOLEAN NOT NULL DEFAULT FALSE, + image_blob VARCHAR(255) REFERENCES blobs(checksum), + favorite boolean not null default FALSE, + CHECK (id != parent_id) +); +CREATE TABLE IF NOT EXISTS saved_filters ( + id serial not null primary key, + name text not null, + mode varchar(255) not null, + find_filter jsonb, + object_filter jsonb, + ui_options jsonb +); +CREATE TABLE IF NOT EXISTS images ( + id serial not null primary key, + title text, + rating smallint, + studio_id integer, + o_counter smallint not null default 0, + organized boolean not null default FALSE, + created_at TIMESTAMP WITH TIME ZONE not null, + updated_at TIMESTAMP WITH TIME ZONE not null, + date date, + code text, + photographer text, + details text, + foreign key(studio_id) references studios(id) on delete SET NULL +); +CREATE TABLE image_urls ( + image_id integer NOT NULL, + position integer NOT NULL, + url varchar(2048) NOT NULL, + foreign key(image_id) references images(id) on delete CASCADE, + PRIMARY KEY(image_id, position, url) +); +CREATE TABLE IF NOT EXISTS galleries ( + id serial not null primary key, + folder_id integer, + title text, + date date, + details text, + studio_id integer, + rating smallint, + organized boolean not null default FALSE, + created_at TIMESTAMP WITH TIME ZONE not null, + updated_at TIMESTAMP WITH TIME ZONE not null, + code text, + photographer text, + foreign key(studio_id) references studios(id) on delete SET NULL, + foreign key(folder_id) references folders(id) on delete SET NULL +); +CREATE TABLE gallery_urls ( + gallery_id integer NOT NULL, + position integer NOT NULL, + url varchar(2048) NOT NULL, + foreign key(gallery_id) references galleries(id) on delete CASCADE, + PRIMARY KEY(gallery_id, position, url) +); +CREATE TABLE IF NOT EXISTS scenes ( + id serial not null primary key, + title text, + details text, + date date, + rating smallint, + studio_id integer, + organized boolean not null default FALSE, + created_at TIMESTAMP WITH TIME ZONE not null, + updated_at TIMESTAMP WITH TIME ZONE not null, + code text, + director text, + resume_time float not null default 0, + play_duration float not null default 0, + cover_blob varchar(255) REFERENCES blobs(checksum), + foreign key(studio_id) references studios(id) on delete SET NULL +); +CREATE TABLE IF NOT EXISTS groups ( + id serial not null primary key, + name text not null, + aliases text, + duration integer, + date date, + rating smallint, + studio_id integer REFERENCES studios(id) ON DELETE SET NULL, + director text, + "description" text, + created_at TIMESTAMP WITH TIME ZONE not null, + updated_at TIMESTAMP WITH TIME ZONE not null, + front_image_blob varchar(255) REFERENCES blobs(checksum), + back_image_blob varchar(255) REFERENCES blobs(checksum) +); +CREATE TABLE IF NOT EXISTS group_urls ( + "group_id" integer NOT NULL, + position integer NOT NULL, + url varchar(2048) NOT NULL, + foreign key("group_id") references "groups"(id) on delete CASCADE, + PRIMARY KEY("group_id", position, url) +); +CREATE TABLE IF NOT EXISTS groups_tags ( + "group_id" integer NOT NULL, + tag_id integer NOT NULL, + foreign key("group_id") references "groups"(id) on delete CASCADE, + foreign key(tag_id) references tags(id) on delete CASCADE, + PRIMARY KEY("group_id", tag_id) +); +CREATE TABLE performer_urls ( + performer_id integer NOT NULL, + position integer NOT NULL, + url varchar(2048) NOT NULL, + foreign key(performer_id) references performers(id) on delete CASCADE, + PRIMARY KEY(performer_id, position, url) +); +CREATE TABLE studios_tags ( + studio_id integer NOT NULL, + tag_id integer NOT NULL, + foreign key(studio_id) references studios(id) on delete CASCADE, + foreign key(tag_id) references tags(id) on delete CASCADE, + PRIMARY KEY(studio_id, tag_id) +); +CREATE TABLE IF NOT EXISTS scenes_view_dates ( + scene_id integer not null, + view_date TIMESTAMP WITH TIME ZONE not null, + foreign key(scene_id) references scenes(id) on delete CASCADE +); +CREATE TABLE IF NOT EXISTS scenes_o_dates ( + scene_id integer not null, + o_date TIMESTAMP WITH TIME ZONE not null, + foreign key(scene_id) references scenes(id) on delete CASCADE +); +CREATE TABLE performer_stash_ids ( + performer_id integer, + endpoint varchar(2048), + stash_id uuid, + foreign key(performer_id) references performers(id) on delete CASCADE +); +CREATE TABLE studio_stash_ids ( + studio_id integer, + endpoint varchar(2048), + stash_id uuid, + foreign key(studio_id) references studios(id) on delete CASCADE +); +CREATE TABLE tags_relations ( + parent_id integer, + child_id integer, + primary key (parent_id, child_id), + foreign key (parent_id) references tags(id) on delete cascade, + foreign key (child_id) references tags(id) on delete cascade +); +CREATE TABLE files_fingerprints ( + file_id integer NOT NULL, + type varchar(255) NOT NULL, + fingerprint text NOT NULL, + foreign key(file_id) references files(id) on delete CASCADE, + PRIMARY KEY (file_id, type, fingerprint) +); +CREATE TABLE video_files ( + file_id integer NOT NULL primary key, + duration float NOT NULL, + video_codec varchar(255) NOT NULL, + format varchar(255) NOT NULL, + audio_codec varchar(255) NOT NULL, + width smallint NOT NULL, + height smallint NOT NULL, + frame_rate float NOT NULL, + bit_rate integer NOT NULL, + interactive boolean not null default FALSE, + interactive_speed int, + foreign key(file_id) references files(id) on delete CASCADE +); +CREATE TABLE video_captions ( + file_id integer NOT NULL, + language_code varchar(255) NOT NULL, + filename varchar(255) NOT NULL, + caption_type varchar(255) NOT NULL, + primary key (file_id, language_code, caption_type), + foreign key(file_id) references video_files(file_id) on delete CASCADE +); +CREATE TABLE image_files ( + file_id integer NOT NULL primary key, + format varchar(255) NOT NULL, + width smallint NOT NULL, + height smallint NOT NULL, + foreign key(file_id) references files(id) on delete CASCADE +); +CREATE TABLE images_files ( + image_id integer NOT NULL, + file_id integer NOT NULL, + "primary" boolean NOT NULL, + foreign key(image_id) references images(id) on delete CASCADE, + foreign key(file_id) references files(id) on delete CASCADE, + PRIMARY KEY(image_id, file_id) +); +CREATE TABLE galleries_files ( + gallery_id integer NOT NULL, + file_id integer NOT NULL, + "primary" boolean NOT NULL, + foreign key(gallery_id) references galleries(id) on delete CASCADE, + foreign key(file_id) references files(id) on delete CASCADE, + PRIMARY KEY(gallery_id, file_id) +); +CREATE TABLE scenes_files ( + scene_id integer NOT NULL, + file_id integer NOT NULL, + "primary" boolean NOT NULL, + foreign key(scene_id) references scenes(id) on delete CASCADE, + foreign key(file_id) references files(id) on delete CASCADE, + PRIMARY KEY(scene_id, file_id) +); +CREATE TABLE IF NOT EXISTS performers_scenes ( + performer_id integer, + scene_id integer, + foreign key(performer_id) references performers(id) on delete CASCADE, + foreign key(scene_id) references scenes(id) on delete CASCADE, + PRIMARY KEY (scene_id, performer_id) +); +CREATE TABLE IF NOT EXISTS scene_markers ( + id serial not null primary key, + title text NOT NULL, + seconds FLOAT NOT NULL, + primary_tag_id INTEGER NOT NULL, + scene_id INTEGER NOT NULL, + created_at TIMESTAMP WITH TIME ZONE NOT NULL, + updated_at TIMESTAMP WITH TIME ZONE NOT NULL, + FOREIGN KEY(primary_tag_id) REFERENCES tags(id), + FOREIGN KEY(scene_id) REFERENCES scenes(id) +); +CREATE TABLE IF NOT EXISTS scene_markers_tags ( + scene_marker_id integer, + tag_id integer, + foreign key(scene_marker_id) references scene_markers(id) on delete CASCADE, + foreign key(tag_id) references tags(id) on delete CASCADE, + PRIMARY KEY(scene_marker_id, tag_id) +); +CREATE TABLE IF NOT EXISTS scenes_tags ( + scene_id integer, + tag_id integer, + foreign key(scene_id) references scenes(id) on delete CASCADE, + foreign key(tag_id) references tags(id) on delete CASCADE, + PRIMARY KEY(scene_id, tag_id) +); +CREATE TABLE IF NOT EXISTS groups_scenes ( + "group_id" integer, + scene_id integer, + scene_index smallint, + foreign key("group_id") references "groups"(id) on delete cascade, + foreign key(scene_id) references scenes(id) on delete cascade, + PRIMARY KEY("group_id", scene_id) +); +CREATE TABLE IF NOT EXISTS performers_images ( + performer_id integer, + image_id integer, + foreign key(performer_id) references performers(id) on delete CASCADE, + foreign key(image_id) references images(id) on delete CASCADE, + PRIMARY KEY(image_id, performer_id) +); +CREATE TABLE IF NOT EXISTS images_tags ( + image_id integer, + tag_id integer, + foreign key(image_id) references images(id) on delete CASCADE, + foreign key(tag_id) references tags(id) on delete CASCADE, + PRIMARY KEY(image_id, tag_id) +); +CREATE TABLE IF NOT EXISTS scene_stash_ids ( + scene_id integer NOT NULL, + endpoint varchar(2048) NOT NULL, + stash_id uuid NOT NULL, + foreign key(scene_id) references scenes(id) on delete CASCADE, + PRIMARY KEY(scene_id, endpoint) +); +CREATE TABLE IF NOT EXISTS scenes_galleries ( + scene_id integer NOT NULL, + gallery_id integer NOT NULL, + foreign key(scene_id) references scenes(id) on delete CASCADE, + foreign key(gallery_id) references galleries(id) on delete CASCADE, + PRIMARY KEY(scene_id, gallery_id) +); +CREATE TABLE IF NOT EXISTS galleries_images ( + gallery_id integer NOT NULL, + image_id integer NOT NULL, + cover boolean not null default FALSE, + foreign key(gallery_id) references galleries(id) on delete CASCADE, + foreign key(image_id) references images(id) on delete CASCADE, + PRIMARY KEY(gallery_id, image_id) +); +CREATE TABLE IF NOT EXISTS performers_galleries ( + performer_id integer NOT NULL, + gallery_id integer NOT NULL, + foreign key(performer_id) references performers(id) on delete CASCADE, + foreign key(gallery_id) references galleries(id) on delete CASCADE, + PRIMARY KEY(gallery_id, performer_id) +); +CREATE TABLE IF NOT EXISTS galleries_tags ( + gallery_id integer NOT NULL, + tag_id integer NOT NULL, + foreign key(gallery_id) references galleries(id) on delete CASCADE, + foreign key(tag_id) references tags(id) on delete CASCADE, + PRIMARY KEY(gallery_id, tag_id) +); +CREATE TABLE IF NOT EXISTS performers_tags ( + performer_id integer NOT NULL, + tag_id integer NOT NULL, + foreign key(performer_id) references performers(id) on delete CASCADE, + foreign key(tag_id) references tags(id) on delete CASCADE, + PRIMARY KEY(performer_id, tag_id) +); +CREATE TABLE IF NOT EXISTS tag_aliases ( + tag_id integer NOT NULL, + alias text NOT NULL, + foreign key(tag_id) references tags(id) on delete CASCADE, + PRIMARY KEY(tag_id, alias) +); +CREATE TABLE IF NOT EXISTS studio_aliases ( + studio_id integer NOT NULL, + alias text NOT NULL, + foreign key(studio_id) references studios(id) on delete CASCADE, + PRIMARY KEY(studio_id, alias) +); +CREATE TABLE performer_aliases ( + performer_id integer NOT NULL, + alias text NOT NULL, + foreign key(performer_id) references performers(id) on delete CASCADE, + PRIMARY KEY(performer_id, alias) +); +CREATE TABLE galleries_chapters ( + id serial not null primary key, + title text not null, + image_index integer not null, + gallery_id integer not null, + created_at TIMESTAMP WITH TIME ZONE not null, + updated_at TIMESTAMP WITH TIME ZONE not null, + foreign key(gallery_id) references galleries(id) on delete CASCADE +); +CREATE TABLE scene_urls ( + scene_id integer NOT NULL, + position integer NOT NULL, + url varchar(2048) NOT NULL, + foreign key(scene_id) references scenes(id) on delete CASCADE, + PRIMARY KEY(scene_id, position, url) +); +CREATE TABLE groups_relations ( + containing_id integer not null, + sub_id integer not null, + order_index integer not null, + description text, + primary key (containing_id, sub_id), + foreign key (containing_id) references groups(id) on delete cascade, + foreign key (sub_id) references groups(id) on delete cascade, + check (containing_id != sub_id) +); +CREATE INDEX index_tags_on_name on tags (name); +CREATE INDEX index_folders_on_parent_folder_id on folders (parent_folder_id); +CREATE UNIQUE INDEX index_folders_on_path_unique on folders (path); +CREATE UNIQUE INDEX index_files_zip_basename_unique ON files (zip_file_id, parent_folder_id, basename) WHERE zip_file_id IS NOT NULL; +CREATE UNIQUE INDEX index_files_on_parent_folder_id_basename_unique on files (parent_folder_id, basename); +CREATE INDEX index_files_on_basename on files (basename); +CREATE INDEX index_folders_on_zip_file_id on folders (zip_file_id) WHERE zip_file_id IS NOT NULL; +CREATE INDEX index_fingerprint_type_fingerprint ON files_fingerprints (type, fingerprint); +CREATE INDEX index_files_fingerprints_file_id ON files_fingerprints (file_id); +CREATE INDEX index_images_files_on_file_id on images_files (file_id); +CREATE UNIQUE INDEX unique_index_images_files_on_primary on images_files (image_id) WHERE "primary" = TRUE; +CREATE INDEX index_galleries_files_file_id ON galleries_files (file_id); +CREATE UNIQUE INDEX unique_index_galleries_files_on_primary on galleries_files (gallery_id) WHERE "primary" = TRUE; +CREATE INDEX index_scenes_files_file_id ON scenes_files (file_id); +CREATE UNIQUE INDEX unique_index_scenes_files_on_primary on scenes_files (scene_id) WHERE "primary" = TRUE; +CREATE INDEX index_performer_stash_ids_on_performer_id ON performer_stash_ids (performer_id); +CREATE INDEX index_studio_stash_ids_on_studio_id ON studio_stash_ids (studio_id); +CREATE INDEX index_performers_scenes_on_performer_id on performers_scenes (performer_id); +CREATE INDEX index_scene_markers_tags_on_tag_id on scene_markers_tags (tag_id); +CREATE INDEX index_scenes_tags_on_tag_id on scenes_tags (tag_id); +CREATE INDEX index_movies_scenes_on_movie_id on groups_scenes (group_id); +CREATE INDEX index_performers_images_on_performer_id on performers_images (performer_id); +CREATE INDEX index_images_tags_on_tag_id on images_tags (tag_id); +CREATE INDEX index_scenes_galleries_on_gallery_id on scenes_galleries (gallery_id); +CREATE INDEX index_galleries_images_on_image_id on galleries_images (image_id); +CREATE INDEX index_performers_galleries_on_performer_id on performers_galleries (performer_id); +CREATE INDEX index_galleries_tags_on_tag_id on galleries_tags (tag_id); +CREATE INDEX index_performers_tags_on_tag_id on performers_tags (tag_id); +CREATE UNIQUE INDEX tag_aliases_alias_unique on tag_aliases (alias); +CREATE UNIQUE INDEX studio_aliases_alias_unique on studio_aliases (alias); +CREATE INDEX performer_aliases_alias on performer_aliases (alias); +CREATE INDEX index_galleries_chapters_on_gallery_id on galleries_chapters (gallery_id); +CREATE INDEX scene_urls_url on scene_urls (url); +CREATE INDEX index_scene_markers_on_primary_tag_id ON scene_markers(primary_tag_id); +CREATE INDEX index_scene_markers_on_scene_id ON scene_markers(scene_id); +CREATE UNIQUE INDEX index_studios_on_name_unique ON studios(name); +CREATE UNIQUE INDEX index_saved_filters_on_mode_name_unique on saved_filters (mode, name); +CREATE INDEX image_urls_url on image_urls (url); +CREATE INDEX index_images_on_studio_id on images (studio_id); +CREATE INDEX gallery_urls_url on gallery_urls (url); +CREATE INDEX index_galleries_on_studio_id on galleries (studio_id); +CREATE UNIQUE INDEX index_galleries_on_folder_id_unique on galleries (folder_id); +CREATE INDEX index_scenes_on_studio_id on scenes (studio_id); +CREATE INDEX performers_urls_url on performer_urls (url); +CREATE UNIQUE INDEX performers_name_disambiguation_unique on performers (name, disambiguation) WHERE disambiguation IS NOT NULL; +CREATE UNIQUE INDEX performers_name_unique on performers (name) WHERE disambiguation IS NULL; +CREATE INDEX index_studios_tags_on_tag_id on studios_tags (tag_id); +CREATE INDEX index_scenes_view_dates ON scenes_view_dates (scene_id); +CREATE INDEX index_scenes_o_dates ON scenes_o_dates (scene_id); +CREATE INDEX index_groups_on_name ON groups(name); +CREATE INDEX index_groups_on_studio_id on groups (studio_id); +CREATE INDEX group_urls_url on group_urls (url); +CREATE INDEX index_groups_tags_on_tag_id on groups_tags (tag_id); +CREATE INDEX index_groups_tags_on_movie_id on groups_tags (group_id); +CREATE UNIQUE INDEX index_galleries_images_gallery_id_cover on galleries_images (gallery_id, cover) WHERE cover = TRUE; +CREATE INDEX index_groups_relations_sub_id ON groups_relations (sub_id); +CREATE UNIQUE INDEX index_groups_relations_order_index_unique ON groups_relations (containing_id, order_index); diff --git a/pkg/postgres/migrations/2_image_studio_index.up.sql b/pkg/postgres/migrations/2_image_studio_index.up.sql new file mode 100644 index 0000000000..3c28cf1196 --- /dev/null +++ b/pkg/postgres/migrations/2_image_studio_index.up.sql @@ -0,0 +1,7 @@ +-- with the existing index, if no images have a studio id, then the index is +-- not used when filtering by studio id. The assumption with this change is that +-- most images don't have a studio id, so filtering by non-null studio id should +-- be faster with this index. This is a tradeoff, as filtering by null studio id +-- will be slower. +DROP INDEX index_images_on_studio_id; +CREATE INDEX index_images_on_studio_id on images (studio_id) WHERE studio_id IS NOT NULL; \ No newline at end of file diff --git a/pkg/postgres/migrations/3_stash_id_updated_at.up.sql b/pkg/postgres/migrations/3_stash_id_updated_at.up.sql new file mode 100644 index 0000000000..8bf9a8cb00 --- /dev/null +++ b/pkg/postgres/migrations/3_stash_id_updated_at.up.sql @@ -0,0 +1,3 @@ +ALTER TABLE performer_stash_ids ADD COLUMN updated_at timestamp not null default '1970-01-01T00:00:00Z'; +ALTER TABLE scene_stash_ids ADD COLUMN updated_at timestamp not null default '1970-01-01T00:00:00Z'; +ALTER TABLE studio_stash_ids ADD COLUMN updated_at timestamp not null default '1970-01-01T00:00:00Z'; diff --git a/pkg/postgres/migrations/4_markers_end.up.sql b/pkg/postgres/migrations/4_markers_end.up.sql new file mode 100644 index 0000000000..05469953ac --- /dev/null +++ b/pkg/postgres/migrations/4_markers_end.up.sql @@ -0,0 +1 @@ +ALTER TABLE scene_markers ADD COLUMN end_seconds FLOAT; \ No newline at end of file diff --git a/pkg/postgres/migrations/5_custom_fields.up.sql b/pkg/postgres/migrations/5_custom_fields.up.sql new file mode 100644 index 0000000000..5855500bae --- /dev/null +++ b/pkg/postgres/migrations/5_custom_fields.up.sql @@ -0,0 +1,10 @@ +-- 71_custom_fields.up.sql +CREATE TABLE performer_custom_fields ( + performer_id integer NOT NULL, + field text NOT NULL, + "value" jsonb NOT NULL, + PRIMARY KEY ("performer_id", "field"), + foreign key("performer_id") references "performers"("id") on delete CASCADE +); + +CREATE INDEX "index_performer_custom_fields_field_value" ON "performer_custom_fields" ("field", "value"); diff --git a/pkg/postgres/migrations/6_tag_sort_name.up.sql b/pkg/postgres/migrations/6_tag_sort_name.up.sql new file mode 100644 index 0000000000..58040f6d27 --- /dev/null +++ b/pkg/postgres/migrations/6_tag_sort_name.up.sql @@ -0,0 +1,2 @@ +-- 72_tag_sort_name.up.sql +ALTER TABLE "tags" ADD COLUMN "sort_name" varchar(255); diff --git a/pkg/postgres/migrations/7_cf_type_field.up.sql b/pkg/postgres/migrations/7_cf_type_field.up.sql new file mode 100644 index 0000000000..1aca695f92 --- /dev/null +++ b/pkg/postgres/migrations/7_cf_type_field.up.sql @@ -0,0 +1,2 @@ +-- 73_cf_type_field.up.sql +ALTER TABLE performer_custom_fields ADD COLUMN "type" text; diff --git a/pkg/postgres/migrations/8_native_basename.up.sql b/pkg/postgres/migrations/8_native_basename.up.sql new file mode 100644 index 0000000000..1c299fc3d6 --- /dev/null +++ b/pkg/postgres/migrations/8_native_basename.up.sql @@ -0,0 +1,17 @@ +CREATE OR REPLACE FUNCTION basename(path text, slash text DEFAULT '/') +RETURNS text AS $$ +DECLARE + re text; + result text; +BEGIN + IF slash IS NULL OR length(slash) != 1 THEN + RAISE EXCEPTION 'Slash must be a single character'; + END IF; + + re := '[^' || regexp_replace(slash, '(.)', '\\\\\1', 'g') || ']+$'; + + result := substring(path FROM re); + + RETURN result; +END; +$$ LANGUAGE plpgsql IMMUTABLE STRICT; diff --git a/pkg/postgres/migrations/9_perl_regex.up.sql b/pkg/postgres/migrations/9_perl_regex.up.sql new file mode 100644 index 0000000000..a34ebd4207 --- /dev/null +++ b/pkg/postgres/migrations/9_perl_regex.up.sql @@ -0,0 +1,6 @@ +CREATE EXTENSION IF NOT EXISTS plperl; +CREATE OR REPLACE FUNCTION regex_match(text, text) +RETURNS boolean AS $$ + my ($string, $pattern) = @_; + return $string =~ /$pattern/u ? 1 : 0; +$$ LANGUAGE plperl IMMUTABLE; diff --git a/pkg/postgres/performer.go b/pkg/postgres/performer.go new file mode 100644 index 0000000000..58a9d9607b --- /dev/null +++ b/pkg/postgres/performer.go @@ -0,0 +1,960 @@ +package postgres + +import ( + "context" + "database/sql" + "errors" + "fmt" + "slices" + + "github.com/doug-martin/goqu/v9" + "github.com/doug-martin/goqu/v9/exp" + "github.com/jmoiron/sqlx" + "github.com/stashapp/stash/pkg/models" + "github.com/stashapp/stash/pkg/utils" + "gopkg.in/guregu/null.v4" + "gopkg.in/guregu/null.v4/zero" +) + +const ( + performerTable = "performers" + performerIDColumn = "performer_id" + performersAliasesTable = "performer_aliases" + performerAliasColumn = "alias" + performersTagsTable = "performers_tags" + + performerURLsTable = "performer_urls" + performerURLColumn = "url" + + performerImageBlobColumn = "image_blob" +) + +type performerRow struct { + ID int `db:"id" goqu:"skipinsert"` + Name null.String `db:"name"` // TODO: make schema non-nullable + Disambigation zero.String `db:"disambiguation"` + Gender zero.String `db:"gender"` + Birthdate NullDate `db:"birthdate"` + BirthdatePrecision null.Int `db:"birthdate_precision"` + Ethnicity zero.String `db:"ethnicity"` + Country zero.String `db:"country"` + EyeColor zero.String `db:"eye_color"` + Height null.Int `db:"height"` + Measurements zero.String `db:"measurements"` + FakeTits zero.String `db:"fake_tits"` + PenisLength null.Float `db:"penis_length"` + Circumcised zero.String `db:"circumcised"` + CareerLength zero.String `db:"career_length"` + Tattoos zero.String `db:"tattoos"` + Piercings zero.String `db:"piercings"` + Favorite bool `db:"favorite"` + CreatedAt Timestamp `db:"created_at"` + UpdatedAt Timestamp `db:"updated_at"` + // expressed as 1-100 + Rating null.Int `db:"rating"` + Details zero.String `db:"details"` + DeathDate NullDate `db:"death_date"` + DeathDatePrecision null.Int `db:"death_date_precision"` + HairColor zero.String `db:"hair_color"` + Weight null.Int `db:"weight"` + IgnoreAutoTag bool `db:"ignore_auto_tag"` + + // not used in resolution or updates + ImageBlob zero.String `db:"image_blob"` +} + +func (r *performerRow) fromPerformer(o models.Performer) { + r.ID = o.ID + r.Name = null.StringFrom(o.Name) + r.Disambigation = zero.StringFrom(o.Disambiguation) + if o.Gender != nil && o.Gender.IsValid() { + r.Gender = zero.StringFrom(o.Gender.String()) + } + r.Birthdate = NullDateFromDatePtr(o.Birthdate) + r.BirthdatePrecision = datePrecisionFromDatePtr(o.Birthdate) + r.Ethnicity = zero.StringFrom(o.Ethnicity) + r.Country = zero.StringFrom(o.Country) + r.EyeColor = zero.StringFrom(o.EyeColor) + r.Height = intFromPtr(o.Height) + r.Measurements = zero.StringFrom(o.Measurements) + r.FakeTits = zero.StringFrom(o.FakeTits) + r.PenisLength = null.FloatFromPtr(o.PenisLength) + if o.Circumcised != nil && o.Circumcised.IsValid() { + r.Circumcised = zero.StringFrom(o.Circumcised.String()) + } + r.CareerLength = zero.StringFrom(o.CareerLength) + r.Tattoos = zero.StringFrom(o.Tattoos) + r.Piercings = zero.StringFrom(o.Piercings) + r.Favorite = o.Favorite + r.CreatedAt = Timestamp{Timestamp: o.CreatedAt} + r.UpdatedAt = Timestamp{Timestamp: o.UpdatedAt} + r.Rating = intFromPtr(o.Rating) + r.Details = zero.StringFrom(o.Details) + r.DeathDate = NullDateFromDatePtr(o.DeathDate) + r.DeathDatePrecision = datePrecisionFromDatePtr(o.DeathDate) + r.HairColor = zero.StringFrom(o.HairColor) + r.Weight = intFromPtr(o.Weight) + r.IgnoreAutoTag = o.IgnoreAutoTag +} + +func (r *performerRow) resolve() *models.Performer { + ret := &models.Performer{ + ID: r.ID, + Name: r.Name.String, + Disambiguation: r.Disambigation.String, + Birthdate: r.Birthdate.DatePtr(r.BirthdatePrecision), + Ethnicity: r.Ethnicity.String, + Country: r.Country.String, + EyeColor: r.EyeColor.String, + Height: nullIntPtr(r.Height), + Measurements: r.Measurements.String, + FakeTits: r.FakeTits.String, + PenisLength: nullFloatPtr(r.PenisLength), + CareerLength: r.CareerLength.String, + Tattoos: r.Tattoos.String, + Piercings: r.Piercings.String, + Favorite: r.Favorite, + CreatedAt: r.CreatedAt.Timestamp.UTC(), + UpdatedAt: r.UpdatedAt.Timestamp.UTC(), + // expressed as 1-100 + Rating: nullIntPtr(r.Rating), + Details: r.Details.String, + DeathDate: r.DeathDate.DatePtr(r.DeathDatePrecision), + HairColor: r.HairColor.String, + Weight: nullIntPtr(r.Weight), + IgnoreAutoTag: r.IgnoreAutoTag, + } + + if r.Gender.ValueOrZero() != "" { + v := models.GenderEnum(r.Gender.String) + ret.Gender = &v + } + + if r.Circumcised.ValueOrZero() != "" { + v := models.CircumisedEnum(r.Circumcised.String) + ret.Circumcised = &v + } + + return ret +} + +type performerRowRecord struct { + updateRecord +} + +func (r *performerRowRecord) fromPartial(o models.PerformerPartial) { + r.setString("name", o.Name) + r.setNullString("disambiguation", o.Disambiguation) + r.setNullString("gender", o.Gender) + r.setNullDate("birthdate", "birthdate_precision", o.Birthdate) + r.setNullString("ethnicity", o.Ethnicity) + r.setNullString("country", o.Country) + r.setNullString("eye_color", o.EyeColor) + r.setNullInt("height", o.Height) + r.setNullString("measurements", o.Measurements) + r.setNullString("fake_tits", o.FakeTits) + r.setNullFloat64("penis_length", o.PenisLength) + r.setNullString("circumcised", o.Circumcised) + r.setNullString("career_length", o.CareerLength) + r.setNullString("tattoos", o.Tattoos) + r.setNullString("piercings", o.Piercings) + r.setBool("favorite", o.Favorite) + r.setTimestamp("created_at", o.CreatedAt) + r.setTimestamp("updated_at", o.UpdatedAt) + r.setNullInt("rating", o.Rating) + r.setNullString("details", o.Details) + r.setNullDate("death_date", "death_date_precision", o.DeathDate) + r.setNullString("hair_color", o.HairColor) + r.setNullInt("weight", o.Weight) + r.setBool("ignore_auto_tag", o.IgnoreAutoTag) +} + +type performerRepositoryType struct { + repository + + tags joinRepository + stashIDs stashIDRepository + + scenes joinRepository + images joinRepository + galleries joinRepository +} + +var ( + performerRepository = performerRepositoryType{ + repository: repository{ + tableName: performerTable, + idColumn: idColumn, + }, + tags: joinRepository{ + repository: repository{ + tableName: performersTagsTable, + idColumn: performerIDColumn, + }, + fkColumn: tagIDColumn, + foreignTable: tagTable, + orderBy: tagTableSortSQL, + }, + stashIDs: stashIDRepository{ + repository{ + tableName: "performer_stash_ids", + idColumn: performerIDColumn, + }, + }, + scenes: joinRepository{ + repository: repository{ + tableName: performersScenesTable, + idColumn: performerIDColumn, + }, + fkColumn: sceneIDColumn, + foreignTable: sceneTable, + }, + images: joinRepository{ + repository: repository{ + tableName: performersImagesTable, + idColumn: performerIDColumn, + }, + fkColumn: imageIDColumn, + foreignTable: imageTable, + }, + galleries: joinRepository{ + repository: repository{ + tableName: performersGalleriesTable, + idColumn: performerIDColumn, + }, + fkColumn: galleryIDColumn, + foreignTable: galleryTable, + }, + } +) + +type PerformerStore struct { + blobJoinQueryBuilder + customFieldsStore + + tableMgr *table +} + +func NewPerformerStore(blobStore *BlobStore) *PerformerStore { + return &PerformerStore{ + blobJoinQueryBuilder: blobJoinQueryBuilder{ + blobStore: blobStore, + joinTable: performerTable, + }, + customFieldsStore: customFieldsStore{ + table: performersCustomFieldsTable, + fk: performersCustomFieldsTable.Col(performerIDColumn), + }, + tableMgr: performerTableMgr, + } +} + +func (qb *PerformerStore) table() exp.IdentifierExpression { + return qb.tableMgr.table +} + +func (qb *PerformerStore) selectDataset() *goqu.SelectDataset { + return dialect.From(qb.table()).Select(qb.table().All()) +} + +func (qb *PerformerStore) Create(ctx context.Context, newObject *models.CreatePerformerInput) error { + var r performerRow + r.fromPerformer(*newObject.Performer) + + id, err := qb.tableMgr.insertID(ctx, r) + if err != nil { + return err + } + + if newObject.Aliases.Loaded() { + if err := performersAliasesTableMgr.insertJoins(ctx, id, newObject.Aliases.List()); err != nil { + return err + } + } + + if newObject.URLs.Loaded() { + const startPos = 0 + if err := performersURLsTableMgr.insertJoins(ctx, id, startPos, newObject.URLs.List()); err != nil { + return err + } + } + + if newObject.TagIDs.Loaded() { + if err := performersTagsTableMgr.insertJoins(ctx, id, newObject.TagIDs.List()); err != nil { + return err + } + } + + if newObject.StashIDs.Loaded() { + if err := performersStashIDsTableMgr.insertJoins(ctx, id, newObject.StashIDs.List()); err != nil { + return err + } + } + + const partial = false + if err := qb.setCustomFields(ctx, id, newObject.CustomFields, partial); err != nil { + return err + } + + updated, err := qb.find(ctx, id) + if err != nil { + return fmt.Errorf("finding after create: %w", err) + } + + *newObject.Performer = *updated + + return nil +} + +func (qb *PerformerStore) UpdatePartial(ctx context.Context, id int, partial models.PerformerPartial) (*models.Performer, error) { + r := performerRowRecord{ + updateRecord{ + Record: make(exp.Record), + }, + } + + r.fromPartial(partial) + + if len(r.Record) > 0 { + if err := qb.tableMgr.updateByID(ctx, id, r.Record); err != nil { + return nil, err + } + } + + if partial.Aliases != nil { + if err := performersAliasesTableMgr.modifyJoins(ctx, id, partial.Aliases.Values, partial.Aliases.Mode); err != nil { + return nil, err + } + } + + if partial.URLs != nil { + if err := performersURLsTableMgr.modifyJoins(ctx, id, partial.URLs.Values, partial.URLs.Mode); err != nil { + return nil, err + } + } + + if partial.TagIDs != nil { + if err := performersTagsTableMgr.modifyJoins(ctx, id, partial.TagIDs.IDs, partial.TagIDs.Mode); err != nil { + return nil, err + } + } + if partial.StashIDs != nil { + if err := performersStashIDsTableMgr.modifyJoins(ctx, id, partial.StashIDs.StashIDs, partial.StashIDs.Mode); err != nil { + return nil, err + } + } + + if err := qb.SetCustomFields(ctx, id, partial.CustomFields); err != nil { + return nil, err + } + + return qb.find(ctx, id) +} + +func (qb *PerformerStore) Update(ctx context.Context, updatedObject *models.UpdatePerformerInput) error { + var r performerRow + r.fromPerformer(*updatedObject.Performer) + + if err := qb.tableMgr.updateByID(ctx, updatedObject.ID, r); err != nil { + return err + } + + if updatedObject.Aliases.Loaded() { + if err := performersAliasesTableMgr.replaceJoins(ctx, updatedObject.ID, updatedObject.Aliases.List()); err != nil { + return err + } + } + + if updatedObject.URLs.Loaded() { + if err := performersURLsTableMgr.replaceJoins(ctx, updatedObject.ID, updatedObject.URLs.List()); err != nil { + return err + } + } + + if updatedObject.TagIDs.Loaded() { + if err := performersTagsTableMgr.replaceJoins(ctx, updatedObject.ID, updatedObject.TagIDs.List()); err != nil { + return err + } + } + + if updatedObject.StashIDs.Loaded() { + if err := performersStashIDsTableMgr.replaceJoins(ctx, updatedObject.ID, updatedObject.StashIDs.List()); err != nil { + return err + } + } + + if err := qb.SetCustomFields(ctx, updatedObject.ID, updatedObject.CustomFields); err != nil { + return err + } + + return nil +} + +func (qb *PerformerStore) Destroy(ctx context.Context, id int) error { + // must handle image checksums manually + if err := qb.destroyImage(ctx, id); err != nil { + return err + } + + return performerRepository.destroyExisting(ctx, []int{id}) +} + +// returns nil, nil if not found +func (qb *PerformerStore) Find(ctx context.Context, id int) (*models.Performer, error) { + ret, err := qb.find(ctx, id) + if errors.Is(err, sql.ErrNoRows) { + return nil, nil + } + return ret, err +} + +func (qb *PerformerStore) FindMany(ctx context.Context, ids []int) ([]*models.Performer, error) { + tableMgr := performerTableMgr + ret := make([]*models.Performer, len(ids)) + + if len(ids) == 0 { + return ret, nil + } + + if err := batchExec(ids, defaultBatchSize, func(batch []int) error { + q := goqu.Select("*").From(tableMgr.table).Where(tableMgr.byIDInts(batch...)) + unsorted, err := qb.getMany(ctx, q) + if err != nil { + return err + } + + for _, s := range unsorted { + i := slices.Index(ids, s.ID) + ret[i] = s + } + + return nil + }); err != nil { + return nil, err + } + + for i := range ret { + if ret[i] == nil { + return nil, fmt.Errorf("performer with id %d not found", ids[i]) + } + } + + return ret, nil +} + +// returns nil, sql.ErrNoRows if not found +func (qb *PerformerStore) find(ctx context.Context, id int) (*models.Performer, error) { + q := qb.selectDataset().Where(qb.tableMgr.byID(id)) + + ret, err := qb.get(ctx, q) + if err != nil { + return nil, err + } + + return ret, nil +} + +func (qb *PerformerStore) findBySubquery(ctx context.Context, sq *goqu.SelectDataset) ([]*models.Performer, error) { + table := qb.table() + + q := qb.selectDataset().Where( + table.Col(idColumn).Eq( + sq, + ), + ) + + return qb.getMany(ctx, q) +} + +// returns nil, sql.ErrNoRows if not found +func (qb *PerformerStore) get(ctx context.Context, q *goqu.SelectDataset) (*models.Performer, error) { + ret, err := qb.getMany(ctx, q) + if err != nil { + return nil, err + } + + if len(ret) == 0 { + return nil, sql.ErrNoRows + } + + return ret[0], nil +} + +func (qb *PerformerStore) getMany(ctx context.Context, q *goqu.SelectDataset) ([]*models.Performer, error) { + const single = false + var ret []*models.Performer + if err := queryFunc(ctx, q, single, func(r *sqlx.Rows) error { + var f performerRow + if err := r.StructScan(&f); err != nil { + return err + } + + s := f.resolve() + + ret = append(ret, s) + return nil + }); err != nil { + return nil, err + } + + return ret, nil +} + +func (qb *PerformerStore) FindBySceneID(ctx context.Context, sceneID int) ([]*models.Performer, error) { + sq := dialect.From(scenesPerformersJoinTable).Select(scenesPerformersJoinTable.Col(performerIDColumn)).Where( + scenesPerformersJoinTable.Col(sceneIDColumn).Eq(sceneID), + ) + ret, err := qb.findBySubquery(ctx, sq) + + if err != nil { + return nil, fmt.Errorf("getting performers for scene %d: %w", sceneID, err) + } + + return ret, nil +} + +func (qb *PerformerStore) FindByImageID(ctx context.Context, imageID int) ([]*models.Performer, error) { + sq := dialect.From(performersImagesJoinTable).Select(performersImagesJoinTable.Col(performerIDColumn)).Where( + performersImagesJoinTable.Col(imageIDColumn).Eq(imageID), + ) + ret, err := qb.findBySubquery(ctx, sq) + + if err != nil { + return nil, fmt.Errorf("getting performers for image %d: %w", imageID, err) + } + + return ret, nil +} + +func (qb *PerformerStore) FindByGalleryID(ctx context.Context, galleryID int) ([]*models.Performer, error) { + sq := dialect.From(performersGalleriesJoinTable).Select(performersGalleriesJoinTable.Col(performerIDColumn)).Where( + performersGalleriesJoinTable.Col(galleryIDColumn).Eq(galleryID), + ) + ret, err := qb.findBySubquery(ctx, sq) + + if err != nil { + return nil, fmt.Errorf("getting performers for gallery %d: %w", galleryID, err) + } + + return ret, nil +} + +func (qb *PerformerStore) FindByNames(ctx context.Context, names []string, nocase bool) ([]*models.Performer, error) { + clause := "name " + if nocase { + clause += "COLLATE NOCASE " + } + clause += "IN " + getInBinding(len(names)) + + var args []interface{} + for _, name := range names { + args = append(args, name) + } + + sq := qb.selectDataset().Prepared(true).Where( + goqu.L(clause, args...), + ) + ret, err := qb.getMany(ctx, sq) + + if err != nil { + return nil, fmt.Errorf("getting performers by names: %w", err) + } + + return ret, nil +} + +func (qb *PerformerStore) CountByTagID(ctx context.Context, tagID int) (int, error) { + joinTable := performersTagsJoinTable + + q := dialect.Select(goqu.COUNT("*")).From(joinTable).Where(joinTable.Col(tagIDColumn).Eq(tagID)) + return count(ctx, q) +} + +func (qb *PerformerStore) Count(ctx context.Context) (int, error) { + q := dialect.Select(goqu.COUNT("*")).From(qb.table()) + return count(ctx, q) +} + +func (qb *PerformerStore) All(ctx context.Context) ([]*models.Performer, error) { + table := qb.table() + return qb.getMany(ctx, qb.selectDataset().Order(table.Col("name").Asc())) +} + +func (qb *PerformerStore) QueryForAutoTag(ctx context.Context, words []string) ([]*models.Performer, error) { + // TODO - Query needs to be changed to support queries of this type, and + // this method should be removed + table := qb.table() + sq := dialect.From(table).Select(table.Col(idColumn)) + // TODO - disabled alias matching until we get finer control over it + // .LeftJoin( + // performersAliasesJoinTable, + // goqu.On(performersAliasesJoinTable.Col(performerIDColumn).Eq(table.Col(idColumn))), + // ) + + var whereClauses []exp.Expression + + for _, w := range words { + whereClauses = append(whereClauses, table.Col("name").ILike(w+"%")) + // TODO - see above + // whereClauses = append(whereClauses, performersAliasesJoinTable.Col("alias").ILike(w+"%")) + } + + sq = sq.Where( + goqu.Or(whereClauses...), + table.Col("ignore_auto_tag").IsFalse(), + ) + + ret, err := qb.findBySubquery(ctx, sq) + + if err != nil { + return nil, fmt.Errorf("getting performers for autotag: %w", err) + } + + return ret, nil +} + +func (qb *PerformerStore) makeQuery(ctx context.Context, performerFilter *models.PerformerFilterType, findFilter *models.FindFilterType) (*queryBuilder, error) { + if performerFilter == nil { + performerFilter = &models.PerformerFilterType{} + } + if findFilter == nil { + findFilter = &models.FindFilterType{} + } + + query := performerRepository.newQuery() + distinctIDs(&query, performerTable) + + if q := findFilter.Q; q != nil && *q != "" { + query.join(performersAliasesTable, "", "performer_aliases.performer_id = performers.id") + searchColumns := []string{"performers.name", "performer_aliases.alias"} + query.parseQueryString(searchColumns, *q) + } + + filter := filterBuilderFromHandler(ctx, &performerFilterHandler{ + performerFilter: performerFilter, + }) + + if err := query.addFilter(filter); err != nil { + return nil, err + } + + var err error + var agg []string + query.sort, agg, err = qb.getPerformerSort(findFilter) + if err != nil { + return nil, err + } + query.pagination = getPagination(findFilter) + query.addGroupBy(agg...) + + return &query, nil +} + +func (qb *PerformerStore) Query(ctx context.Context, performerFilter *models.PerformerFilterType, findFilter *models.FindFilterType) ([]*models.Performer, int, error) { + query, err := qb.makeQuery(ctx, performerFilter, findFilter) + if err != nil { + return nil, 0, err + } + + idsResult, countResult, err := query.executeFind(ctx) + if err != nil { + return nil, 0, err + } + + performers, err := qb.FindMany(ctx, idsResult) + if err != nil { + return nil, 0, err + } + + return performers, countResult, nil +} + +func (qb *PerformerStore) QueryCount(ctx context.Context, performerFilter *models.PerformerFilterType, findFilter *models.FindFilterType) (int, error) { + query, err := qb.makeQuery(ctx, performerFilter, findFilter) + if err != nil { + return 0, err + } + + return query.executeCount(ctx) +} + +func (qb *PerformerStore) sortByOCounter(direction string) string { + // need to sum the o_counter from scenes and images + return " ORDER BY (" + selectPerformerOCountSQL + ") " + direction +} + +func (qb *PerformerStore) sortByPlayCount(direction string) string { + // need to sum the o_counter from scenes and images + return " ORDER BY (" + selectPerformerPlayCountSQL + ") " + direction +} + +// used for sorting on performer last o_date +var selectPerformerLastOAtSQL = utils.StrFormat( + "SELECT MAX(o_date) FROM ("+ + "SELECT {o_date} FROM {performers_scenes} s "+ + "LEFT JOIN {scenes} ON {scenes}.id = s.{scene_id} "+ + "LEFT JOIN {scenes_o_dates} ON {scenes_o_dates}.{scene_id} = {scenes}.id "+ + "WHERE s.{performer_id} = {performers}.id"+ + ")", + map[string]interface{}{ + "performer_id": performerIDColumn, + "performers": performerTable, + "performers_scenes": performersScenesTable, + "scenes": sceneTable, + "scene_id": sceneIDColumn, + "scenes_o_dates": scenesODatesTable, + "o_date": sceneODateColumn, + }, +) + +func (qb *PerformerStore) sortByLastOAt(direction string) string { + // need to get the o_dates from scenes + return " ORDER BY (" + selectPerformerLastOAtSQL + ") " + direction +} + +// used for sorting on performer last view_date +var selectPerformerLastPlayedAtSQL = utils.StrFormat( + "SELECT MAX(view_date) FROM ("+ + "SELECT {view_date} FROM {performers_scenes} s "+ + "LEFT JOIN {scenes} ON {scenes}.id = s.{scene_id} "+ + "LEFT JOIN {scenes_view_dates} ON {scenes_view_dates}.{scene_id} = {scenes}.id "+ + "WHERE s.{performer_id} = {performers}.id"+ + ")", + map[string]interface{}{ + "performer_id": performerIDColumn, + "performers": performerTable, + "performers_scenes": performersScenesTable, + "scenes": sceneTable, + "scene_id": sceneIDColumn, + "scenes_view_dates": scenesViewDatesTable, + "view_date": sceneViewDateColumn, + }, +) + +func (qb *PerformerStore) sortByLastPlayedAt(direction string) string { + // need to get the view_dates from scenes + return " ORDER BY (" + selectPerformerLastPlayedAtSQL + ") " + direction +} + +// used for sorting by total scene duration +var selectPerformerScenesDurationSQL = utils.StrFormat( + "SELECT COALESCE(SUM(video_files.duration), 0) FROM {performers_scenes} s "+ + "LEFT JOIN {scenes} ON {scenes}.id = s.{scene_id} "+ + "LEFT JOIN {scenes_files} ON {scenes_files}.{scene_id} = {scenes}.id "+ + "LEFT JOIN video_files ON video_files.file_id = {scenes_files}.file_id "+ + "WHERE s.{performer_id} = {performers}.id", + map[string]interface{}{ + "performer_id": performerIDColumn, + "performers": performerTable, + "performers_scenes": performersScenesTable, + "scenes": sceneTable, + "scene_id": sceneIDColumn, + "scenes_files": scenesFilesTable, + }, +) + +func (qb *PerformerStore) sortByScenesDuration(direction string) string { + // need to sum duration from all scenes for this performer + return " ORDER BY (" + selectPerformerScenesDurationSQL + ") " + direction +} + +var performerSortOptions = sortOptions{ + "birthdate", + "career_length", + "created_at", + "galleries_count", + "height", + "id", + "images_count", + "last_o_at", + "last_played_at", + "measurements", + "name", + "o_counter", + "penis_length", + "play_count", + "random", + "rating", + "scenes_count", + "scenes_duration", + "tag_count", + "updated_at", + "weight", +} + +func (qb *PerformerStore) getPerformerSort(findFilter *models.FindFilterType) (string, []string, error) { + var sort string + var direction string + if findFilter == nil { + sort = "name" + direction = "ASC" + } else { + sort = findFilter.GetSort("name") + direction = findFilter.GetDirection() + } + + // CVE-2024-32231 - ensure sort is in the list of allowed sorts + if err := performerSortOptions.validateSort(sort); err != nil { + return "", nil, err + } + + var agg []string + sortQuery := "" + switch sort { + case "tag_count": + sortQuery += getCountSort(performerTable, performersTagsTable, performerIDColumn, direction) + case "scenes_count": + sortQuery += getCountSort(performerTable, performersScenesTable, performerIDColumn, direction) + case "scenes_duration": + sortQuery += qb.sortByScenesDuration(direction) + case "images_count": + sortQuery += getCountSort(performerTable, performersImagesTable, performerIDColumn, direction) + case "galleries_count": + sortQuery += getCountSort(performerTable, performersGalleriesTable, performerIDColumn, direction) + case "play_count": + sortQuery += qb.sortByPlayCount(direction) + case "o_counter": + sortQuery += qb.sortByOCounter(direction) + case "last_played_at": + sortQuery += qb.sortByLastPlayedAt(direction) + case "last_o_at": + sortQuery += qb.sortByLastOAt(direction) + default: + var add string + add, agg = getSort(sort, direction, "performers") + sortQuery += add + } + + // Whatever the sorting, always use name/id as a final sort + sortQuery += ", COALESCE(performers.name, CAST(performers.id as text)) COLLATE NATURAL_CI ASC" + agg = append(agg, "performers.name", "performers.id") + return sortQuery, agg, nil +} + +func (qb *PerformerStore) GetTagIDs(ctx context.Context, id int) ([]int, error) { + return performerRepository.tags.getIDs(ctx, id) +} + +func (qb *PerformerStore) GetImage(ctx context.Context, performerID int) ([]byte, error) { + return qb.blobJoinQueryBuilder.GetImage(ctx, performerID, performerImageBlobColumn) +} + +func (qb *PerformerStore) HasImage(ctx context.Context, performerID int) (bool, error) { + return qb.blobJoinQueryBuilder.HasImage(ctx, performerID, performerImageBlobColumn) +} + +func (qb *PerformerStore) UpdateImage(ctx context.Context, performerID int, image []byte) error { + return qb.blobJoinQueryBuilder.UpdateImage(ctx, performerID, performerImageBlobColumn, image) +} + +func (qb *PerformerStore) destroyImage(ctx context.Context, performerID int) error { + return qb.blobJoinQueryBuilder.DestroyImage(ctx, performerID, performerImageBlobColumn) +} + +func (qb *PerformerStore) GetAliases(ctx context.Context, performerID int) ([]string, error) { + return performersAliasesTableMgr.get(ctx, performerID) +} + +func (qb *PerformerStore) GetURLs(ctx context.Context, performerID int) ([]string, error) { + return performersURLsTableMgr.get(ctx, performerID) +} + +func (qb *PerformerStore) GetStashIDs(ctx context.Context, performerID int) ([]models.StashID, error) { + return performersStashIDsTableMgr.get(ctx, performerID) +} + +func (qb *PerformerStore) FindByStashID(ctx context.Context, stashID models.StashID) ([]*models.Performer, error) { + sq := dialect.From(performersStashIDsJoinTable).Select(performersStashIDsJoinTable.Col(performerIDColumn)).Where( + performersStashIDsJoinTable.Col("stash_id").Eq(stashID.StashID), + performersStashIDsJoinTable.Col("endpoint").Eq(stashID.Endpoint), + ) + ret, err := qb.findBySubquery(ctx, sq) + + if err != nil { + return nil, fmt.Errorf("getting performers for stash ID %s: %w", stashID.StashID, err) + } + + return ret, nil +} + +func (qb *PerformerStore) FindByStashIDStatus(ctx context.Context, hasStashID bool, stashboxEndpoint string) ([]*models.Performer, error) { + table := qb.table() + sq := dialect.From(table).LeftJoin( + performersStashIDsJoinTable, + goqu.On(table.Col(idColumn).Eq(performersStashIDsJoinTable.Col(performerIDColumn))), + ).Select(table.Col(idColumn)) + + if hasStashID { + sq = sq.Where( + performersStashIDsJoinTable.Col("stash_id").IsNotNull(), + performersStashIDsJoinTable.Col("endpoint").Eq(stashboxEndpoint), + ) + } else { + sq = sq.Where( + performersStashIDsJoinTable.Col("stash_id").IsNull(), + ) + } + + ret, err := qb.findBySubquery(ctx, sq) + + if err != nil { + return nil, fmt.Errorf("getting performers for stash-box endpoint %s: %w", stashboxEndpoint, err) + } + + return ret, nil +} + +func (qb *PerformerStore) Merge(ctx context.Context, source []int, destination int) error { + if len(source) == 0 { + return nil + } + + inBinding := getInBinding(len(source)) + + args := []interface{}{destination} + srcArgs := make([]interface{}, len(source)) + for i, id := range source { + if id == destination { + return errors.New("cannot merge where source == destination") + } + srcArgs[i] = id + } + + args = append(args, srcArgs...) + + performerTables := map[string]string{ + performersScenesTable: sceneIDColumn, + performersGalleriesTable: galleryIDColumn, + performersImagesTable: imageIDColumn, + performersTagsTable: tagIDColumn, + } + + args = append(args, destination) + + // for each table, update source performer ids to destination performer id, ignoring duplicates + for table, idColumn := range performerTables { + _, err := dbWrapper.Exec(ctx, `UPDATE OR IGNORE `+table+` +SET performer_id = ? +WHERE performer_id IN `+inBinding+` +AND NOT EXISTS(SELECT 1 FROM `+table+` o WHERE o.`+idColumn+` = `+table+`.`+idColumn+` AND o.performer_id = ?)`, + args..., + ) + if err != nil { + return err + } + + // delete source performer ids from the table where they couldn't be set + if _, err := dbWrapper.Exec(ctx, `DELETE FROM `+table+` WHERE performer_id IN `+inBinding, srcArgs...); err != nil { + return err + } + } + + for _, id := range source { + err := qb.Destroy(ctx, id) + if err != nil { + return err + } + } + + return nil +} diff --git a/pkg/postgres/performer_filter.go b/pkg/postgres/performer_filter.go new file mode 100644 index 0000000000..3c285afd34 --- /dev/null +++ b/pkg/postgres/performer_filter.go @@ -0,0 +1,698 @@ +package postgres + +import ( + "context" + "fmt" + "strconv" + "strings" + + "github.com/stashapp/stash/pkg/models" + "github.com/stashapp/stash/pkg/utils" +) + +type performerFilterHandler struct { + performerFilter *models.PerformerFilterType +} + +func (qb *performerFilterHandler) validate() error { + filter := qb.performerFilter + if filter == nil { + return nil + } + + if err := validateFilterCombination(filter.OperatorFilter); err != nil { + return err + } + + if subFilter := filter.SubFilter(); subFilter != nil { + sqb := &performerFilterHandler{performerFilter: subFilter} + if err := sqb.validate(); err != nil { + return err + } + } + + // if legacy height filter used, ensure only supported modifiers are used + if filter.Height != nil { + // treat as an int filter + intCrit := &models.IntCriterionInput{ + Modifier: filter.Height.Modifier, + } + if !intCrit.ValidModifier() { + return fmt.Errorf("invalid height modifier: %s", filter.Height.Modifier) + } + + // ensure value is a valid number + if _, err := strconv.Atoi(filter.Height.Value); err != nil { + return fmt.Errorf("invalid height value: %s", filter.Height.Value) + } + } + + return nil +} + +func (qb *performerFilterHandler) handle(ctx context.Context, f *filterBuilder) { + filter := qb.performerFilter + if filter == nil { + return + } + + if err := qb.validate(); err != nil { + f.setError(err) + return + } + + sf := filter.SubFilter() + if sf != nil { + sub := &performerFilterHandler{sf} + handleSubFilter(ctx, sub, f, filter.OperatorFilter) + } + + f.handleCriterion(ctx, qb.criterionHandler()) +} + +func (qb *performerFilterHandler) criterionHandler() criterionHandler { + filter := qb.performerFilter + const tableName = performerTable + heightCmCrit := filter.HeightCm + + return compoundHandler{ + stringCriterionHandler(filter.Name, tableName+".name"), + stringCriterionHandler(filter.Disambiguation, tableName+".disambiguation"), + stringCriterionHandler(filter.Details, tableName+".details"), + + boolCriterionHandler(filter.FilterFavorites, tableName+".favorite", nil), + boolCriterionHandler(filter.IgnoreAutoTag, tableName+".ignore_auto_tag", nil), + + yearFilterCriterionHandler(filter.BirthYear, tableName+".birthdate"), + yearFilterCriterionHandler(filter.DeathYear, tableName+".death_date"), + + qb.performerAgeFilterCriterionHandler(filter.Age), + + criterionHandlerFunc(func(ctx context.Context, f *filterBuilder) { + if gender := filter.Gender; gender != nil { + genderCopy := *gender + if genderCopy.Value.IsValid() && len(genderCopy.ValueList) == 0 { + genderCopy.ValueList = []models.GenderEnum{genderCopy.Value} + } + + v := utils.StringerSliceToStringSlice(genderCopy.ValueList) + enumCriterionHandler(genderCopy.Modifier, v, tableName+".gender")(ctx, f) + } + }), + + qb.performerIsMissingCriterionHandler(filter.IsMissing), + stringCriterionHandler(filter.Ethnicity, tableName+".ethnicity"), + stringCriterionHandler(filter.Country, tableName+".country"), + stringCriterionHandler(filter.EyeColor, tableName+".eye_color"), + + // special handler for legacy height filter + criterionHandlerFunc(func(ctx context.Context, f *filterBuilder) { + if heightCmCrit == nil && filter.Height != nil { + heightCm, _ := strconv.Atoi(filter.Height.Value) // already validated + heightCmCrit = &models.IntCriterionInput{ + Value: heightCm, + Modifier: filter.Height.Modifier, + } + } + }), + + intCriterionHandler(heightCmCrit, tableName+".height", nil), + + stringCriterionHandler(filter.Measurements, tableName+".measurements"), + stringCriterionHandler(filter.FakeTits, tableName+".fake_tits"), + floatCriterionHandler(filter.PenisLength, tableName+".penis_length", nil), + + criterionHandlerFunc(func(ctx context.Context, f *filterBuilder) { + if circumcised := filter.Circumcised; circumcised != nil { + v := utils.StringerSliceToStringSlice(circumcised.Value) + enumCriterionHandler(circumcised.Modifier, v, tableName+".circumcised")(ctx, f) + } + }), + + stringCriterionHandler(filter.CareerLength, tableName+".career_length"), + stringCriterionHandler(filter.Tattoos, tableName+".tattoos"), + stringCriterionHandler(filter.Piercings, tableName+".piercings"), + intCriterionHandler(filter.Rating100, tableName+".rating", nil), + stringCriterionHandler(filter.HairColor, tableName+".hair_color"), + qb.urlsCriterionHandler(filter.URL), + intCriterionHandler(filter.Weight, tableName+".weight", nil), + criterionHandlerFunc(func(ctx context.Context, f *filterBuilder) { + if filter.StashID != nil { + performerRepository.stashIDs.join(f, "performer_stash_ids", "performers.id") + stringCriterionHandler(filter.StashID, "performer_stash_ids.stash_id")(ctx, f) + } + }), + &stashIDCriterionHandler{ + c: filter.StashIDEndpoint, + stashIDRepository: &performerRepository.stashIDs, + stashIDTableAs: "performer_stash_ids", + parentIDCol: "performers.id", + }, + &stashIDsCriterionHandler{ + c: filter.StashIDsEndpoint, + stashIDRepository: &performerRepository.stashIDs, + stashIDTableAs: "performer_stash_ids", + parentIDCol: "performers.id", + }, + + qb.aliasCriterionHandler(filter.Aliases), + + qb.tagsCriterionHandler(filter.Tags), + + qb.studiosCriterionHandler(filter.Studios), + + qb.groupsCriterionHandler(filter.Groups), + + qb.appearsWithCriterionHandler(filter.Performers), + + qb.tagCountCriterionHandler(filter.TagCount), + qb.sceneCountCriterionHandler(filter.SceneCount), + qb.imageCountCriterionHandler(filter.ImageCount), + qb.galleryCountCriterionHandler(filter.GalleryCount), + qb.playCounterCriterionHandler(filter.PlayCount), + qb.oCounterCriterionHandler(filter.OCounter), + &dateCriterionHandler{filter.Birthdate, tableName + ".birthdate", nil}, + &dateCriterionHandler{filter.DeathDate, tableName + ".death_date", nil}, + ×tampCriterionHandler{filter.CreatedAt, tableName + ".created_at", nil}, + ×tampCriterionHandler{filter.UpdatedAt, tableName + ".updated_at", nil}, + + &relatedFilterHandler{ + relatedIDCol: "performers_scenes.scene_id", + relatedRepo: sceneRepository.repository, + relatedHandler: &sceneFilterHandler{filter.ScenesFilter}, + joinFn: func(f *filterBuilder) { + performerRepository.scenes.innerJoin(f, "", "performers.id") + }, + }, + + &relatedFilterHandler{ + relatedIDCol: "performers_images.image_id", + relatedRepo: imageRepository.repository, + relatedHandler: &imageFilterHandler{filter.ImagesFilter}, + joinFn: func(f *filterBuilder) { + performerRepository.images.innerJoin(f, "", "performers.id") + }, + }, + + &relatedFilterHandler{ + relatedIDCol: "performers_galleries.gallery_id", + relatedRepo: galleryRepository.repository, + relatedHandler: &galleryFilterHandler{filter.GalleriesFilter}, + joinFn: func(f *filterBuilder) { + performerRepository.galleries.innerJoin(f, "", "performers.id") + }, + }, + + &relatedFilterHandler{ + relatedIDCol: "performer_tag.tag_id", + relatedRepo: tagRepository.repository, + relatedHandler: &tagFilterHandler{filter.TagsFilter}, + joinFn: func(f *filterBuilder) { + performerRepository.tags.innerJoin(f, "performer_tag", "performers.id") + }, + }, + + &customFieldsFilterHandler{ + table: performersCustomFieldsTable.GetTable(), + fkCol: performerIDColumn, + c: filter.CustomFields, + idCol: "performers.id", + }, + } +} + +// TODO - we need to provide a whitelist of possible values +func (qb *performerFilterHandler) performerIsMissingCriterionHandler(isMissing *string) criterionHandlerFunc { + return func(ctx context.Context, f *filterBuilder) { + if isMissing != nil && *isMissing != "" { + switch *isMissing { + case "url": + performersURLsTableMgr.join(f, "", "performers.id") + f.addWhere("performer_urls.url IS NULL") + case "scenes": // Deprecated: use `scene_count == 0` filter instead + f.addLeftJoin(performersScenesTable, "scenes_join", "scenes_join.performer_id = performers.id") + f.addWhere("scenes_join.scene_id IS NULL") + case "image": + f.addWhere("performers.image_blob IS NULL") + case "stash_id": + performersStashIDsTableMgr.join(f, "performer_stash_ids", "performers.id") + f.addWhere("performer_stash_ids.performer_id IS NULL") + case "aliases": + performersAliasesTableMgr.join(f, "", "performers.id") + f.addWhere("performer_aliases.alias IS NULL") + default: + f.addWhere("(performers." + *isMissing + " IS NULL OR TRIM(CAST(performers." + *isMissing + " AS TEXT)) = '')") + } + } + } +} + +func (qb *performerFilterHandler) performerAgeFilterCriterionHandler(age *models.IntCriterionInput) criterionHandlerFunc { + return func(ctx context.Context, f *filterBuilder) { + if age != nil && age.Modifier.IsValid() { + clause, args := getIntCriterionWhereClause( + "EXTRACT(YEAR FROM AGE(COALESCE(performers.death_date, CURRENT_DATE), performers.birthdate))", + *age, + ) + f.addWhere(clause, args...) + } + } +} + +func (qb *performerFilterHandler) urlsCriterionHandler(url *models.StringCriterionInput) criterionHandlerFunc { + h := stringListCriterionHandlerBuilder{ + primaryTable: performerTable, + primaryFK: performerIDColumn, + joinTable: performerURLsTable, + stringColumn: performerURLColumn, + addJoinTable: func(f *filterBuilder) { + performersURLsTableMgr.join(f, "", "performers.id") + }, + } + + return h.handler(url) +} + +func (qb *performerFilterHandler) aliasCriterionHandler(alias *models.StringCriterionInput) criterionHandlerFunc { + h := stringListCriterionHandlerBuilder{ + primaryTable: performerTable, + primaryFK: performerIDColumn, + joinTable: performersAliasesTable, + stringColumn: performerAliasColumn, + addJoinTable: func(f *filterBuilder) { + performersAliasesTableMgr.join(f, "", "performers.id") + }, + } + + return h.handler(alias) +} + +func (qb *performerFilterHandler) tagsCriterionHandler(tags *models.HierarchicalMultiCriterionInput) criterionHandlerFunc { + h := joinedHierarchicalMultiCriterionHandlerBuilder{ + primaryTable: performerTable, + foreignTable: tagTable, + foreignFK: "tag_id", + + relationsTable: "tags_relations", + joinAs: "performer_tag", + joinTable: performersTagsTable, + primaryFK: performerIDColumn, + } + + return h.handler(tags) +} + +func (qb *performerFilterHandler) tagCountCriterionHandler(count *models.IntCriterionInput) criterionHandlerFunc { + h := countCriterionHandlerBuilder{ + primaryTable: performerTable, + joinTable: performersTagsTable, + primaryFK: performerIDColumn, + } + + return h.handler(count) +} + +func (qb *performerFilterHandler) sceneCountCriterionHandler(count *models.IntCriterionInput) criterionHandlerFunc { + h := countCriterionHandlerBuilder{ + primaryTable: performerTable, + joinTable: performersScenesTable, + primaryFK: performerIDColumn, + } + + return h.handler(count) +} + +func (qb *performerFilterHandler) imageCountCriterionHandler(count *models.IntCriterionInput) criterionHandlerFunc { + h := countCriterionHandlerBuilder{ + primaryTable: performerTable, + joinTable: performersImagesTable, + primaryFK: performerIDColumn, + } + + return h.handler(count) +} + +func (qb *performerFilterHandler) galleryCountCriterionHandler(count *models.IntCriterionInput) criterionHandlerFunc { + h := countCriterionHandlerBuilder{ + primaryTable: performerTable, + joinTable: performersGalleriesTable, + primaryFK: performerIDColumn, + } + + return h.handler(count) +} + +// used for sorting and filtering on performer o-count +var selectPerformerOCountSQL = utils.StrFormat( + "SELECT SUM(o_counter) "+ + "FROM ("+ + "SELECT SUM(o_counter) as o_counter from {performers_images} s "+ + "LEFT JOIN {images} ON {images}.id = s.{images_id} "+ + "WHERE s.{performer_id} = {performers}.id "+ + "UNION ALL "+ + "SELECT COUNT({scenes_o_dates}.{o_date}) as o_counter from {performers_scenes} s "+ + "LEFT JOIN {scenes} ON {scenes}.id = s.{scene_id} "+ + "LEFT JOIN {scenes_o_dates} ON {scenes_o_dates}.{scene_id} = {scenes}.id "+ + "WHERE s.{performer_id} = {performers}.id "+ + ")", + map[string]interface{}{ + "performers_images": performersImagesTable, + "images": imageTable, + "performer_id": performerIDColumn, + "images_id": imageIDColumn, + "performers": performerTable, + "performers_scenes": performersScenesTable, + "scenes": sceneTable, + "scene_id": sceneIDColumn, + "scenes_o_dates": scenesODatesTable, + "o_date": sceneODateColumn, + }, +) + +// used for sorting and filtering play count on performer view count +var selectPerformerPlayCountSQL = utils.StrFormat( + "SELECT COUNT(DISTINCT {view_date}) FROM ("+ + "SELECT {view_date} FROM {performers_scenes} s "+ + "LEFT JOIN {scenes} ON {scenes}.id = s.{scene_id} "+ + "LEFT JOIN {scenes_view_dates} ON {scenes_view_dates}.{scene_id} = {scenes}.id "+ + "WHERE s.{performer_id} = {performers}.id"+ + ")", + map[string]interface{}{ + "performer_id": performerIDColumn, + "performers": performerTable, + "performers_scenes": performersScenesTable, + "scenes": sceneTable, + "scene_id": sceneIDColumn, + "scenes_view_dates": scenesViewDatesTable, + "view_date": sceneViewDateColumn, + }, +) + +func (qb *performerFilterHandler) oCounterCriterionHandler(count *models.IntCriterionInput) criterionHandlerFunc { + return func(ctx context.Context, f *filterBuilder) { + if count == nil { + return + } + + lhs := "(" + selectPerformerOCountSQL + ")" + clause, args := getIntCriterionWhereClause(lhs, *count) + + f.addWhere(clause, args...) + } +} + +func (qb *performerFilterHandler) playCounterCriterionHandler(count *models.IntCriterionInput) criterionHandlerFunc { + return func(ctx context.Context, f *filterBuilder) { + if count == nil { + return + } + + lhs := "(" + selectPerformerPlayCountSQL + ")" + clause, args := getIntCriterionWhereClause(lhs, *count) + + f.addWhere(clause, args...) + } +} + +func (qb *performerFilterHandler) studiosCriterionHandler(studios *models.HierarchicalMultiCriterionInput) criterionHandlerFunc { + return func(ctx context.Context, f *filterBuilder) { + if studios != nil { + formatMaps := []utils.StrFormatMap{ + { + "primaryTable": sceneTable, + "joinTable": performersScenesTable, + "primaryFK": sceneIDColumn, + }, + { + "primaryTable": imageTable, + "joinTable": performersImagesTable, + "primaryFK": imageIDColumn, + }, + { + "primaryTable": galleryTable, + "joinTable": performersGalleriesTable, + "primaryFK": galleryIDColumn, + }, + } + + if studios.Modifier == models.CriterionModifierIsNull || studios.Modifier == models.CriterionModifierNotNull { + var notClause string + if studios.Modifier == models.CriterionModifierNotNull { + notClause = "NOT" + } + + var conditions []string + for _, c := range formatMaps { + f.addLeftJoin(c["joinTable"].(string), "", fmt.Sprintf("%s.performer_id = performers.id", c["joinTable"])) + f.addLeftJoin(c["primaryTable"].(string), "", fmt.Sprintf("%s.%s = %s.id", c["joinTable"], c["primaryFK"], c["primaryTable"])) + + conditions = append(conditions, fmt.Sprintf("%s.studio_id IS NULL", c["primaryTable"])) + } + + f.addWhere(fmt.Sprintf("%s (%s)", notClause, strings.Join(conditions, " AND "))) + return + } + + if len(studios.Value) == 0 && len(studios.Excludes) == 0 { + return + } + + var clauseCondition string + + switch studios.Modifier { + case models.CriterionModifierIncludes: + // return performers who appear in scenes/images/galleries with any of the given studios + clauseCondition = "NOT" + case models.CriterionModifierExcludes: + // exclude performers who appear in scenes/images/galleries with any of the given studios + clauseCondition = "" + default: + return + } + + if len(studios.Value) > 0 { + const derivedPerformerStudioTable = "performer_studio" + valuesClause, err := getHierarchicalValues(ctx, studios.Value, studioTable, "", "parent_id", "child_id", studios.Depth, false) + if err != nil { + f.setError(err) + return + } + f.addWith("studio(root_id, item_id) AS (" + valuesClause + ")") + + templStr := `SELECT performer_id FROM {primaryTable} + INNER JOIN {joinTable} ON {primaryTable}.id = {joinTable}.{primaryFK} + INNER JOIN studio ON {primaryTable}.studio_id = studio.item_id` + + var unions []string + for _, c := range formatMaps { + unions = append(unions, utils.StrFormat(templStr, c)) + } + + f.addWith(fmt.Sprintf("%s AS (%s)", derivedPerformerStudioTable, strings.Join(unions, " UNION "))) + + f.addLeftJoin(derivedPerformerStudioTable, "", fmt.Sprintf("performers.id = %s.performer_id", derivedPerformerStudioTable)) + f.addWhere(fmt.Sprintf("%s.performer_id IS %s NULL", derivedPerformerStudioTable, clauseCondition)) + } + + // #6412 - handle excludes as well + if len(studios.Excludes) > 0 { + excludeValuesClause, err := getHierarchicalValues(ctx, studios.Excludes, studioTable, "", "parent_id", "child_id", studios.Depth, false) + if err != nil { + f.setError(err) + return + } + f.addWith("exclude_studio(root_id, item_id) AS (" + excludeValuesClause + ")") + + excludeTemplStr := `SELECT performer_id FROM {primaryTable} + INNER JOIN {joinTable} ON {primaryTable}.id = {joinTable}.{primaryFK} + INNER JOIN exclude_studio ON {primaryTable}.studio_id = exclude_studio.item_id` + + var unions []string + for _, c := range formatMaps { + unions = append(unions, utils.StrFormat(excludeTemplStr, c)) + } + + const excludePerformerStudioTable = "performer_studio_exclude" + f.addWith(fmt.Sprintf("%s AS (%s)", excludePerformerStudioTable, strings.Join(unions, " UNION "))) + + f.addLeftJoin(excludePerformerStudioTable, "", fmt.Sprintf("performers.id = %s.performer_id", excludePerformerStudioTable)) + f.addWhere(fmt.Sprintf("%s.performer_id IS NULL", excludePerformerStudioTable)) + } + } + } +} + +func (qb *performerFilterHandler) groupsCriterionHandler(groups *models.HierarchicalMultiCriterionInput) criterionHandlerFunc { + return func(ctx context.Context, f *filterBuilder) { + if groups != nil { + if groups.Modifier == models.CriterionModifierIsNull || groups.Modifier == models.CriterionModifierNotNull { + var notClause string + if groups.Modifier == models.CriterionModifierNotNull { + notClause = "NOT" + } + + f.addLeftJoin(performersScenesTable, "", "performers_scenes.performer_id = performers.id") + f.addLeftJoin(groupsScenesTable, "", "performers_scenes.scene_id = groups_scenes.scene_id") + + f.addWhere(fmt.Sprintf("%s groups_scenes.group_id IS NULL", notClause)) + return + } + + if len(groups.Value) == 0 { + return + } + + var clauseCondition string + + switch groups.Modifier { + case models.CriterionModifierIncludes: + // return performers who appear in scenes with any of the given groups + clauseCondition = "NOT" + case models.CriterionModifierExcludes: + // exclude performers who appear in scenes with any of the given groups + clauseCondition = "" + default: + return + } + + const derivedPerformerGroupTable = "performer_group" + + // Simplified approach: direct group-scene-performer relationship without hierarchy + var valuesClauses []string + for _, value := range groups.Value { + id, err := strconv.Atoi(value) + if err != nil { + return + } + + valuesClauses = append(valuesClauses, fmt.Sprintf("%d", id)) + } + valuesClause := strings.Join(valuesClauses, ",") + + // If depth is specified and not 0, we need hierarchy, otherwise use simple approach + depthVal := 0 + if groups.Depth != nil { + depthVal = *groups.Depth + } + + if depthVal == 0 { + // Simple case: no hierarchy, direct group relationship + f.addWith("group_values(id) AS (VALUES (" + valuesClause + "))") + + // f.addWith(fmt.Sprintf("group_values(id) AS (VALUES %s)", strings.Repeat("(?),", len(groups.Value)-1)+"(?)"), args...) + + templStr := `SELECT performer_id FROM {joinTable} + INNER JOIN {primaryTable} ON {joinTable}.scene_id = {primaryTable}.scene_id + INNER JOIN group_values ON {primaryTable}.{groupFK} = group_values.id` + + formatMaps := []utils.StrFormatMap{ + { + "primaryTable": groupsScenesTable, + "joinTable": performersScenesTable, + "primaryFK": sceneIDColumn, + "groupFK": groupIDColumn, + }, + } + + var unions []string + for _, c := range formatMaps { + unions = append(unions, utils.StrFormat(templStr, c)) + } + + f.addWith(fmt.Sprintf("%s AS (%s)", derivedPerformerGroupTable, strings.Join(unions, " UNION "))) + } else { + // Complex case: with hierarchy + var depthCondition string + if depthVal != -1 { + depthCondition = fmt.Sprintf("WHERE depth < %d", depthVal) + } + + // Build recursive CTE for group hierarchy + hierarchyQuery := fmt.Sprintf(`group_hierarchy AS ( +SELECT sub_id AS root_id, sub_id AS item_id, 0 AS depth FROM groups_relations WHERE sub_id IN (%s) +UNION +SELECT root_id, sub_id, depth + 1 FROM groups_relations INNER JOIN group_hierarchy ON item_id = containing_id %s +)`, valuesClause, depthCondition) + + f.addRecursiveWith(hierarchyQuery) + + templStr := `SELECT performer_id FROM {joinTable} + INNER JOIN {primaryTable} ON {joinTable}.scene_id = {primaryTable}.scene_id + INNER JOIN group_hierarchy ON {primaryTable}.{groupFK} = group_hierarchy.item_id` + + formatMaps := []utils.StrFormatMap{ + { + "primaryTable": groupsScenesTable, + "joinTable": performersScenesTable, + "primaryFK": sceneIDColumn, + "groupFK": groupIDColumn, + }, + } + + var unions []string + for _, c := range formatMaps { + unions = append(unions, utils.StrFormat(templStr, c)) + } + + f.addWith(fmt.Sprintf("%s AS (%s)", derivedPerformerGroupTable, strings.Join(unions, " UNION "))) + } + + f.addLeftJoin(derivedPerformerGroupTable, "", fmt.Sprintf("performers.id = %s.performer_id", derivedPerformerGroupTable)) + f.addWhere(fmt.Sprintf("%s.performer_id IS %s NULL", derivedPerformerGroupTable, clauseCondition)) + } + } +} + +func (qb *performerFilterHandler) appearsWithCriterionHandler(performers *models.MultiCriterionInput) criterionHandlerFunc { + return func(ctx context.Context, f *filterBuilder) { + if performers != nil { + formatMaps := []utils.StrFormatMap{ + { + "primaryTable": performersScenesTable, + "joinTable": performersScenesTable, + "primaryFK": sceneIDColumn, + }, + { + "primaryTable": performersImagesTable, + "joinTable": performersImagesTable, + "primaryFK": imageIDColumn, + }, + { + "primaryTable": performersGalleriesTable, + "joinTable": performersGalleriesTable, + "primaryFK": galleryIDColumn, + }, + } + + if len(performers.Value) == '0' { + return + } + + const derivedPerformerPerformersTable = "performer_performers" + + valuesClause := strings.Join(performers.Value, "),(") + + f.addWith("performer(id) AS (VALUES(" + valuesClause + "))") + + templStr := `SELECT {primaryTable}2.performer_id FROM {primaryTable} + INNER JOIN {primaryTable} AS {primaryTable}2 ON {primaryTable}.{primaryFK} = {primaryTable}2.{primaryFK} + INNER JOIN performer ON {primaryTable}.performer_id = performer.id + WHERE {primaryTable}2.performer_id != performer.id` + + if performers.Modifier == models.CriterionModifierIncludesAll && len(performers.Value) > 1 { + templStr += ` + GROUP BY {primaryTable}2.performer_id + HAVING(count(distinct {primaryTable}.performer_id) = ` + strconv.Itoa(len(performers.Value)) + `)` + } + + var unions []string + for _, c := range formatMaps { + unions = append(unions, utils.StrFormat(templStr, c)) + } + + f.addWith(fmt.Sprintf("%s AS (%s)", derivedPerformerPerformersTable, strings.Join(unions, " UNION "))) + + f.addInnerJoin(derivedPerformerPerformersTable, "", fmt.Sprintf("performers.id = %s.performer_id", derivedPerformerPerformersTable)) + } + } +} diff --git a/pkg/postgres/query.go b/pkg/postgres/query.go new file mode 100644 index 0000000000..4ca2bca823 --- /dev/null +++ b/pkg/postgres/query.go @@ -0,0 +1,254 @@ +package postgres + +import ( + "context" + "fmt" + "strings" + + "github.com/stashapp/stash/pkg/models" + "github.com/stashapp/stash/pkg/sliceutil" +) + +type queryPagination struct { + page int + perPage int +} + +type queryBuilder struct { + repository *repository + + columns []string + from string + + joins joins + whereClauses []string + havingClauses []string + args []interface{} + withClauses []string + recursiveWith bool + groupByClauses []string + + sort string + pagination *queryPagination +} + +func (qb queryBuilder) body(includeSortPagination bool) string { + return fmt.Sprintf("SELECT %s FROM %s%s", strings.Join(qb.columns, ", "), qb.from, qb.joins.toSQL(includeSortPagination)) +} + +func (qb *queryBuilder) addColumn(column string) { + qb.columns = append(qb.columns, column) +} + +func (qb *queryBuilder) addGroupBy(columns ...string) { + if len(columns) > 0 { + qb.groupByClauses = sliceutil.AppendUniques(qb.groupByClauses, columns) + } +} + +func (qb queryBuilder) toSQL(includeSortPagination bool) string { + body := qb.body(includeSortPagination) + + withClause := "" + if len(qb.withClauses) > 0 { + var recursive string + if qb.recursiveWith { + recursive = " RECURSIVE " + } + withClause = "WITH " + recursive + strings.Join(qb.withClauses, ", ") + " " + } + + body = withClause + qb.repository.buildQueryBody(body, qb.whereClauses, qb.havingClauses, qb.groupByClauses) + if includeSortPagination { + body += qb.sort + if qb.pagination != nil { + body += getPaginationSQL(qb.pagination) + } + } + + return body +} + +func (qb queryBuilder) findIDs(ctx context.Context) ([]int, error) { + const includeSortPagination = true + sql := qb.toSQL(includeSortPagination) + return qb.repository.runIdsQuery(ctx, sql, qb.args) +} + +func (qb queryBuilder) executeFind(ctx context.Context) ([]int, int, error) { + const includeSortPagination = true + body := qb.body(includeSortPagination) + + // Redirect + if qb.pagination != nil && qb.pagination.perPage == 0 { + res, err := qb.executeCount(ctx) + return []int{}, res, err + } + + pagination := getPaginationSQL(qb.pagination) + return qb.repository.executeFindQuery(ctx, body, qb.args, qb.sort, pagination, qb.whereClauses, qb.havingClauses, qb.withClauses, qb.groupByClauses, qb.recursiveWith) +} + +func (qb queryBuilder) executeCount(ctx context.Context) (int, error) { + const includeSortPagination = false + body := qb.body(includeSortPagination) + + withClause := "" + if len(qb.withClauses) > 0 { + var recursive string + if qb.recursiveWith { + recursive = " RECURSIVE " + } + withClause = "WITH " + recursive + strings.Join(qb.withClauses, ", ") + " " + } + + body = qb.repository.buildQueryBody(body, qb.whereClauses, qb.havingClauses, qb.groupByClauses) + countQuery := withClause + qb.repository.buildCountQuery(body) + return qb.repository.runCountQuery(ctx, countQuery, qb.args) +} + +func (qb *queryBuilder) addWhere(clauses ...string) { + for _, clause := range clauses { + if len(clause) > 0 { + qb.whereClauses = append(qb.whereClauses, clause) + } + } +} + +func (qb *queryBuilder) addHaving(clauses ...string) { + for _, clause := range clauses { + if len(clause) > 0 { + qb.havingClauses = append(qb.havingClauses, clause) + } + } +} + +func (qb *queryBuilder) addWith(recursive bool, clauses ...string) { + for _, clause := range clauses { + if len(clause) > 0 { + qb.withClauses = append(qb.withClauses, clause) + } + } + + qb.recursiveWith = qb.recursiveWith || recursive +} + +func (qb *queryBuilder) addArg(args ...interface{}) { + qb.args = append(qb.args, args...) +} + +func (qb *queryBuilder) hasJoin(alias string) bool { + for _, j := range qb.joins { + if j.alias() == alias { + return true + } + } + + return false +} + +func (qb *queryBuilder) join(table, as, onClause string) { + newJoin := join{ + table: table, + as: as, + onClause: onClause, + joinType: "LEFT", + } + + qb.joins.add(newJoin) +} + +func (qb *queryBuilder) joinSort(table, as, onClause string) { + newJoin := join{ + sort: false, // BUG: If we use this, we need to remake getSort since we do groupby on the args. + table: table, + as: as, + onClause: onClause, + joinType: "LEFT", + } + + qb.joins.add(newJoin) +} + +func (qb *queryBuilder) addJoins(joins ...join) { + for _, j := range joins { + if qb.joins.addUnique(j) { + qb.args = append(qb.args, j.args...) + } + } +} + +func (qb *queryBuilder) addFilter(f *filterBuilder) error { + err := f.getError() + if err != nil { + return err + } + + clause, args := f.generateWithClauses() + if len(clause) > 0 { + qb.addWith(f.recursiveWith, clause) + } + + if len(args) > 0 { + // WITH clause always comes first and thus precedes alk args + qb.args = append(args, qb.args...) + } + + // add joins here to insert args + qb.addJoins(f.getAllJoins()...) + + clause, args = f.generateWhereClauses() + if len(clause) > 0 { + qb.addWhere(clause) + } + + if len(args) > 0 { + qb.addArg(args...) + } + + clause, args = f.generateHavingClauses() + if len(clause) > 0 { + qb.addHaving(clause) + } + + if len(args) > 0 { + qb.addArg(args...) + } + + return nil +} + +func (qb *queryBuilder) parseQueryString(columns []string, q string) { + specs := models.ParseSearchString(q) + + for _, t := range specs.MustHave { + var clauses []string + + for _, column := range columns { + clauses = append(clauses, column+" ILIKE ?") + qb.addArg(like(t)) + } + + qb.addWhere("(" + strings.Join(clauses, " OR ") + ")") + } + + for _, t := range specs.MustNot { + for _, column := range columns { + qb.addWhere(coalesce(column) + " NOT ILIKE ?") + qb.addArg(like(t)) + } + } + + for _, set := range specs.AnySets { + var clauses []string + + for _, column := range columns { + for _, v := range set { + clauses = append(clauses, column+" ILIKE ?") + qb.addArg(like(v)) + } + } + + qb.addWhere("(" + strings.Join(clauses, " OR ") + ")") + } +} diff --git a/pkg/postgres/record.go b/pkg/postgres/record.go new file mode 100644 index 0000000000..9b2fdcd74b --- /dev/null +++ b/pkg/postgres/record.go @@ -0,0 +1,108 @@ +package postgres + +import ( + "github.com/doug-martin/goqu/v9/exp" + "github.com/stashapp/stash/pkg/models" + "gopkg.in/guregu/null.v4" + "gopkg.in/guregu/null.v4/zero" +) + +type updateRecord struct { + exp.Record +} + +func (r *updateRecord) set(destField string, v interface{}) { + r.Record[destField] = v +} + +func (r *updateRecord) setString(destField string, v models.OptionalString) { + if v.Set { + if v.Null { + panic("null value not allowed in optional string") + } + r.set(destField, v.Value) + } +} + +func (r *updateRecord) setNullString(destField string, v models.OptionalString) { + if v.Set { + r.set(destField, zero.StringFromPtr(v.Ptr())) + } +} + +func (r *updateRecord) setBool(destField string, v models.OptionalBool) { + if v.Set { + if v.Null { + panic("null value not allowed in optional bool") + } + r.set(destField, v.Value) + } +} + +func (r *updateRecord) setInt(destField string, v models.OptionalInt) { + if v.Set { + if v.Null { + panic("null value not allowed in optional int") + } + r.set(destField, v.Value) + } +} + +func (r *updateRecord) setNullInt(destField string, v models.OptionalInt) { + if v.Set { + r.set(destField, intFromPtr(v.Ptr())) + } +} + +// func (r *updateRecord) setInt64(destField string, v models.OptionalInt64) { +// if v.Set { +// if v.Null { +// panic("null value not allowed in optional int64") +// } +// r.set(destField, v.Value) +// } +// } + +// func (r *updateRecord) setNullInt64(destField string, v models.OptionalInt64) { +// if v.Set { +// r.set(destField, null.IntFromPtr(v.Ptr())) +// } +// } + +func (r *updateRecord) setFloat64(destField string, v models.OptionalFloat64) { + if v.Set { + if v.Null { + panic("null value not allowed in optional float64") + } + r.set(destField, v.Value) + } +} + +func (r *updateRecord) setNullFloat64(destField string, v models.OptionalFloat64) { + if v.Set { + r.set(destField, null.FloatFromPtr(v.Ptr())) + } +} + +func (r *updateRecord) setTimestamp(destField string, v models.OptionalTime) { + if v.Set { + if v.Null { + panic("null value not allowed in optional time") + } + r.set(destField, Timestamp{Timestamp: v.Value}) + } +} + +//nolint:golint,unused +func (r *updateRecord) setNullTimestamp(destField string, v models.OptionalTime) { + if v.Set { + r.set(destField, NullTimestampFromTimePtr(v.Ptr())) + } +} + +func (r *updateRecord) setNullDate(destField string, precisionField string, v models.OptionalDate) { + if v.Set { + r.set(destField, NullDateFromDatePtr(v.Ptr())) + r.set(precisionField, datePrecisionFromDatePtr(v.Ptr())) + } +} diff --git a/pkg/postgres/relationships.go b/pkg/postgres/relationships.go new file mode 100644 index 0000000000..a7a8050e4a --- /dev/null +++ b/pkg/postgres/relationships.go @@ -0,0 +1,41 @@ +package postgres + +import ( + "context" + + "github.com/stashapp/stash/pkg/models" +) + +type idRelationshipStore struct { + joinTable *joinTable +} + +func (s *idRelationshipStore) createRelationships(ctx context.Context, id int, fkIDs models.RelatedIDs) error { + if fkIDs.Loaded() { + if err := s.joinTable.insertJoins(ctx, id, fkIDs.List()); err != nil { + return err + } + } + + return nil +} + +func (s *idRelationshipStore) modifyRelationships(ctx context.Context, id int, fkIDs *models.UpdateIDs) error { + if fkIDs != nil { + if err := s.joinTable.modifyJoins(ctx, id, fkIDs.IDs, fkIDs.Mode); err != nil { + return err + } + } + + return nil +} + +func (s *idRelationshipStore) replaceRelationships(ctx context.Context, id int, fkIDs models.RelatedIDs) error { + if fkIDs.Loaded() { + if err := s.joinTable.replaceJoins(ctx, id, fkIDs.List()); err != nil { + return err + } + } + + return nil +} diff --git a/pkg/postgres/repository.go b/pkg/postgres/repository.go new file mode 100644 index 0000000000..b488aa7074 --- /dev/null +++ b/pkg/postgres/repository.go @@ -0,0 +1,578 @@ +package postgres + +import ( + "context" + "database/sql" + "errors" + "fmt" + "strings" + + "github.com/jackc/pgx/v5/pgconn" + "github.com/jmoiron/sqlx" + + "github.com/stashapp/stash/pkg/models" +) + +const idColumn = "id" + +type repository struct { + tableName string + idColumn string +} + +func (r *repository) getAll(ctx context.Context, id int, f func(rows *sqlx.Rows) error) error { + stmt := fmt.Sprintf("SELECT * FROM %s WHERE %s = ?", r.tableName, r.idColumn) + return r.queryFunc(ctx, stmt, []interface{}{id}, false, f) +} + +func (r *repository) destroyExisting(ctx context.Context, ids []int) error { + for _, id := range ids { + exists, err := r.exists(ctx, id) + if err != nil { + return err + } + + if !exists { + return fmt.Errorf("%s %d does not exist in %s", r.idColumn, id, r.tableName) + } + } + + return r.destroy(ctx, ids) +} + +func (r *repository) destroy(ctx context.Context, ids []int) error { + for _, id := range ids { + stmt := fmt.Sprintf("DELETE FROM %s WHERE %s = ?", r.tableName, r.idColumn) + if _, err := dbWrapper.Exec(ctx, stmt, id); err != nil { + return err + } + } + + return nil +} + +func (r *repository) exists(ctx context.Context, id int) (bool, error) { + stmt := fmt.Sprintf("SELECT %s FROM %s WHERE %s = ? LIMIT 1", r.idColumn, r.tableName, r.idColumn) + stmt = r.buildCountQuery(stmt) + + c, err := r.runCountQuery(ctx, stmt, []interface{}{id}) + if err != nil { + return false, err + } + + return c == 1, nil +} + +func (r *repository) buildCountQuery(query string) string { + return "SELECT COUNT(*) as count FROM (" + query + ") as temp" +} + +func (r *repository) runCountQuery(ctx context.Context, query string, args []interface{}) (int, error) { + result := struct { + Int int `db:"count"` + }{0} + + // Perform query and fetch result + if err := dbWrapper.Get(ctx, &result, query, args...); err != nil && !errors.Is(err, sql.ErrNoRows) { + return 0, err + } + + return result.Int, nil +} + +func (r *repository) runIdsQuery(ctx context.Context, query string, args []interface{}) ([]int, error) { + var result []struct { + Int int `db:"id"` + } + + if err := dbWrapper.Select(ctx, &result, query, args...); err != nil && !errors.Is(err, sql.ErrNoRows) { + return []int{}, fmt.Errorf("running query: %s [%v]: %w", query, args, err) + } + + vsm := make([]int, len(result)) + for i, v := range result { + vsm[i] = v.Int + } + return vsm, nil +} + +type rawRowsWithCount struct { + ID int `db:"id"` + TotalCount sql.NullInt64 `db:"total_count"` + TotalSize sql.NullFloat64 `db:"total_size"` + TotalCustom sql.NullFloat64 `db:"total_custom"` // Misc selectors +} + +type RowsWithCounts struct { + IDs []int + TotalCount sql.NullInt64 + TotalSize sql.NullFloat64 + TotalCustom sql.NullFloat64 +} + +func (r *repository) runIdsWithCount(ctx context.Context, query string, args []interface{}) (*RowsWithCounts, error) { + var result []rawRowsWithCount + + if err := dbWrapper.Select(ctx, &result, query, args...); err != nil && !errors.Is(err, sql.ErrNoRows) { + return nil, fmt.Errorf("running query: %s [%v]: %w", query, args, err) + } + + rows := RowsWithCounts{} + if len(result) == 0 { + return &rows, nil + } + + rows.IDs = make([]int, len(result)) + for i, row := range result { + rows.IDs[i] = row.ID + } + + rows.TotalCount = result[0].TotalCount + rows.TotalSize = result[0].TotalSize + rows.TotalCustom = result[0].TotalCustom + + return &rows, nil +} + +func (r *repository) queryFunc(ctx context.Context, query string, args []interface{}, single bool, f func(rows *sqlx.Rows) error) error { + rows, err := dbWrapper.QueryxContext(ctx, query, args...) + + if err != nil && !errors.Is(err, sql.ErrNoRows) { + return err + } + defer rows.Close() + + for rows.Next() { + if err := f(rows); err != nil { + return err + } + if single { + break + } + } + + if err := rows.Err(); err != nil { + return err + } + + return nil +} + +// queryStruct executes a query and scans the result into the provided struct. +// Unlike the other query methods, this will return an error if no rows are found. +func (r *repository) queryStruct(ctx context.Context, query string, args []interface{}, out interface{}) error { + // changed from queryFunc, since it was not logging the performance correctly, + // since the query doesn't actually execute until Scan is called + if err := dbWrapper.Get(ctx, out, query, args...); err != nil { + return fmt.Errorf("executing query: %s [%v]: %w", query, args, err) + } + + return nil +} + +func (r *repository) querySimple(ctx context.Context, query string, args []interface{}, out interface{}) error { + rows, err := dbWrapper.Queryx(ctx, query, args...) + + if err != nil && !errors.Is(err, sql.ErrNoRows) { + return err + } + defer rows.Close() + + if rows.Next() { + if err := rows.Scan(out); err != nil { + return err + } + } + + if err := rows.Err(); err != nil { + return err + } + + return nil +} + +func (r *repository) buildQueryBody(body string, whereClauses []string, havingClauses []string, groupByClauses []string) string { + if len(whereClauses) > 0 { + body = body + " WHERE " + strings.Join(whereClauses, " AND ") // TODO handle AND or OR + } + if len(havingClauses) > 0 { + groupByClauses = append(groupByClauses, r.tableName+".id") + } + if len(groupByClauses) > 0 { + body += " GROUP BY " + strings.Join(groupByClauses, ", ") + " " + } + if len(havingClauses) > 0 { + body = body + " HAVING " + strings.Join(havingClauses, " AND ") // TODO handle AND or OR + } + + return body +} + +func (r *repository) buildCombinedQuery(body, sort, pagination string) string { + return ` + WITH base_query AS ( + ` + body + ` + ` + sort + ` + ), + total AS ( + SELECT COUNT(*) AS total_count FROM base_query + ) + SELECT base_query.id, total.total_count + FROM base_query, total + ` + pagination +} + +func (r *repository) executeFindQuery(ctx context.Context, body string, args []interface{}, sort string, pagination string, whereClauses []string, havingClauses []string, withClauses []string, groupByClauses []string, recursiveWith bool) ([]int, int, error) { + body = r.buildQueryBody(body, whereClauses, havingClauses, groupByClauses) + + withClause := "" + if len(withClauses) > 0 { + var recursive string + if recursiveWith { + recursive = " RECURSIVE " + } + withClause = "WITH " + recursive + strings.Join(withClauses, ", ") + " " + } + + // Perform query and fetch result + combinedQuery := r.buildCombinedQuery(withClause+body, sort, pagination) + + var obj *RowsWithCounts + var err error + if obj, err = r.runIdsWithCount(ctx, combinedQuery, args); err != nil { + return nil, 0, err + } + + return obj.IDs, int(obj.TotalCount.Int64), err +} + +func (r *repository) newQuery() queryBuilder { + return queryBuilder{ + repository: r, + } +} + +func (r *repository) join(j joiner, as string, parentIDCol string) { + t := r.tableName + if as != "" { + t = as + } + j.addLeftJoin(r.tableName, as, fmt.Sprintf("%s.%s = %s", t, r.idColumn, parentIDCol)) +} + +func (r *repository) innerJoin(j joiner, as string, parentIDCol string) { + t := r.tableName + if as != "" { + t = as + } + j.addInnerJoin(r.tableName, as, fmt.Sprintf("%s.%s = %s", t, r.idColumn, parentIDCol)) +} + +type joiner interface { + addLeftJoin(table, as, onClause string, args ...interface{}) + addInnerJoin(table, as, onClause string, args ...interface{}) +} + +type joinRepository struct { + repository + fkColumn string + + // fields for ordering + foreignTable string + orderBy string +} + +func (r *joinRepository) getIDs(ctx context.Context, id int) ([]int, error) { + var joinStr string + if r.foreignTable != "" { + joinStr = fmt.Sprintf(" INNER JOIN %s ON %[1]s.id = %s.%s", r.foreignTable, r.tableName, r.fkColumn) + } + + query := fmt.Sprintf(`SELECT %[2]s.%[1]s as id from %s%s WHERE %s = ?`, r.fkColumn, r.tableName, joinStr, r.idColumn) + + if r.orderBy != "" { + query += " ORDER BY " + r.orderBy + } + + return r.runIdsQuery(ctx, query, []interface{}{id}) +} + +func (r *joinRepository) insert(ctx context.Context, id int, foreignIDs ...int) error { + stmt, err := dbWrapper.Prepare(ctx, fmt.Sprintf("INSERT INTO %s (%s, %s) VALUES (?, ?)", r.tableName, r.idColumn, r.fkColumn)) + if err != nil { + return err + } + + defer stmt.Close() + + for _, fk := range foreignIDs { + if _, err := dbWrapper.ExecStmt(ctx, stmt, id, fk); err != nil { + return err + } + } + return nil +} + +// insertOrIgnore inserts a join into the table, silently failing in the event that a conflict occurs (ie when the join already exists) +func (r *joinRepository) insertOrIgnore(ctx context.Context, id int, foreignIDs ...int) error { + stmt, err := dbWrapper.Prepare(ctx, fmt.Sprintf("INSERT INTO %s (%s, %s) VALUES (?, ?) ON CONFLICT (%[2]s, %s) DO NOTHING", r.tableName, r.idColumn, r.fkColumn)) + if err != nil { + return err + } + + defer stmt.Close() + + for _, fk := range foreignIDs { + if _, err := dbWrapper.ExecStmt(ctx, stmt, id, fk); err != nil { + return err + } + } + return nil +} + +func (r *joinRepository) destroyJoins(ctx context.Context, id int, foreignIDs ...int) error { + stmt := fmt.Sprintf("DELETE FROM %s WHERE %s = ? AND %s IN %s", r.tableName, r.idColumn, r.fkColumn, getInBinding(len(foreignIDs))) + + args := make([]interface{}, len(foreignIDs)+1) + args[0] = id + for i, v := range foreignIDs { + args[i+1] = v + } + + if _, err := dbWrapper.Exec(ctx, stmt, args...); err != nil { + return err + } + + return nil +} + +func (r *joinRepository) replace(ctx context.Context, id int, foreignIDs []int) error { + if err := r.destroy(ctx, []int{id}); err != nil { + return err + } + + for _, fk := range foreignIDs { + if err := r.insert(ctx, id, fk); err != nil { + return err + } + } + + return nil +} + +type captionRepository struct { + repository +} + +func (r *captionRepository) get(ctx context.Context, id models.FileID) ([]*models.VideoCaption, error) { + query := fmt.Sprintf("SELECT %s, %s, %s from %s WHERE %s = ?", captionCodeColumn, captionFilenameColumn, captionTypeColumn, r.tableName, r.idColumn) + var ret []*models.VideoCaption + err := r.queryFunc(ctx, query, []interface{}{id}, false, func(rows *sqlx.Rows) error { + var captionCode string + var captionFilename string + var captionType string + + if err := rows.Scan(&captionCode, &captionFilename, &captionType); err != nil { + return err + } + + caption := &models.VideoCaption{ + LanguageCode: captionCode, + Filename: captionFilename, + CaptionType: captionType, + } + ret = append(ret, caption) + return nil + }) + return ret, err +} + +func (r *captionRepository) insert(ctx context.Context, id models.FileID, caption *models.VideoCaption) (sql.Result, error) { + stmt := fmt.Sprintf("INSERT INTO %s (%s, %s, %s, %s) VALUES (?, ?, ?, ?)", r.tableName, r.idColumn, captionCodeColumn, captionFilenameColumn, captionTypeColumn) + return dbWrapper.Exec(ctx, stmt, id, caption.LanguageCode, caption.Filename, caption.CaptionType) +} + +func (r *captionRepository) replace(ctx context.Context, id models.FileID, captions []*models.VideoCaption) error { + if err := r.destroy(ctx, []int{int(id)}); err != nil { + return err + } + + for _, caption := range captions { + if _, err := r.insert(ctx, id, caption); err != nil { + return err + } + } + + return nil +} + +type stringRepository struct { + repository + stringColumn string +} + +func (r *stringRepository) get(ctx context.Context, id int) ([]string, error) { + query := fmt.Sprintf("SELECT %s from %s WHERE %s = ?", r.stringColumn, r.tableName, r.idColumn) + var ret []string + err := r.queryFunc(ctx, query, []interface{}{id}, false, func(rows *sqlx.Rows) error { + var out string + if err := rows.Scan(&out); err != nil { + return err + } + + ret = append(ret, out) + return nil + }) + return ret, err +} + +func (r *stringRepository) insert(ctx context.Context, id int, s string) (sql.Result, error) { + stmt := fmt.Sprintf("INSERT INTO %s (%s, %s) VALUES (?, ?)", r.tableName, r.idColumn, r.stringColumn) + return dbWrapper.Exec(ctx, stmt, id, s) +} + +func (r *stringRepository) replace(ctx context.Context, id int, newStrings []string) error { + if err := r.destroy(ctx, []int{id}); err != nil { + return err + } + + for _, s := range newStrings { + if _, err := r.insert(ctx, id, s); err != nil { + return err + } + } + + return nil +} + +type stashIDRepository struct { + repository +} + +type stashIDs []models.StashID + +func (s *stashIDs) Append(o interface{}) { + *s = append(*s, o.(models.StashID)) +} + +func (s *stashIDs) New() interface{} { + return &models.StashID{} +} + +func (r *stashIDRepository) get(ctx context.Context, id int) ([]models.StashID, error) { + query := fmt.Sprintf("SELECT stash_id, endpoint, updated_at from %s WHERE %s = ?", r.tableName, r.idColumn) + var ret stashIDs + err := r.queryFunc(ctx, query, []interface{}{id}, false, func(rows *sqlx.Rows) error { + var v stashIDRow + if err := rows.StructScan(&v); err != nil { + return err + } + ret.Append(v.resolve()) + return nil + }) + return ret, err +} + +type filesRepository struct { + repository +} + +type relatedFileRow struct { + ID int `db:"id"` + FileID models.FileID `db:"file_id"` + Primary bool `db:"primary"` +} + +func idToIndexMap(ids []int) map[int]int { + ret := make(map[int]int) + for i, id := range ids { + ret[id] = i + } + return ret +} + +func (r *filesRepository) getMany(ctx context.Context, ids []int, primaryOnly bool) ([][]models.FileID, error) { + var primaryClause string + if primaryOnly { + primaryClause = ` AND "primary" = true` + } + + query := fmt.Sprintf(`SELECT %s as id, file_id, "primary" from %s WHERE %[1]s IN %[3]s%s`, r.idColumn, r.tableName, getInBinding(len(ids)), primaryClause) + + idi := make([]interface{}, len(ids)) + for i, id := range ids { + idi[i] = id + } + + var fileRows []relatedFileRow + if err := r.queryFunc(ctx, query, idi, false, func(rows *sqlx.Rows) error { + var f relatedFileRow + + if err := rows.StructScan(&f); err != nil { + return err + } + + fileRows = append(fileRows, f) + + return nil + }); err != nil { + return nil, err + } + + ret := make([][]models.FileID, len(ids)) + idToIndex := idToIndexMap(ids) + + for _, row := range fileRows { + id := row.ID + fileID := row.FileID + + if row.Primary { + // prepend to list + ret[idToIndex[id]] = append([]models.FileID{fileID}, ret[idToIndex[id]]...) + } else { + ret[idToIndex[id]] = append(ret[idToIndex[id]], row.FileID) + } + } + + return ret, nil +} + +func (r *filesRepository) get(ctx context.Context, id int) ([]models.FileID, error) { + query := fmt.Sprintf(`SELECT file_id, "primary" from %s WHERE %s = ?`, r.tableName, r.idColumn) + + type relatedFile struct { + FileID models.FileID `db:"file_id"` + Primary bool `db:"primary"` + } + + var ret []models.FileID + if err := r.queryFunc(ctx, query, []interface{}{id}, false, func(rows *sqlx.Rows) error { + var f relatedFile + + if err := rows.StructScan(&f); err != nil { + return err + } + + if f.Primary { + // prepend to list + ret = append([]models.FileID{f.FileID}, ret...) + } else { + ret = append(ret, f.FileID) + } + + return nil + }); err != nil { + return nil, err + } + + return ret, nil +} + +func (r *repository) isConstraintError(err error) bool { + var pgErr *pgconn.PgError + if errors.As(err, &pgErr) { + // Class 23 — Integrity Constraint Violation + return pgErr.Code[:2] == "23" + } + return false +} diff --git a/pkg/postgres/saved_filter.go b/pkg/postgres/saved_filter.go new file mode 100644 index 0000000000..224f93e0ab --- /dev/null +++ b/pkg/postgres/saved_filter.go @@ -0,0 +1,261 @@ +package postgres + +import ( + "context" + "database/sql" + "encoding/json" + "errors" + "fmt" + "slices" + + "github.com/doug-martin/goqu/v9" + "github.com/doug-martin/goqu/v9/exp" + "github.com/jmoiron/sqlx" + + "github.com/stashapp/stash/pkg/logger" + "github.com/stashapp/stash/pkg/models" +) + +const ( + savedFilterTable = "saved_filters" + savedFilterDefaultName = "" +) + +type savedFilterRow struct { + ID int `db:"id" goqu:"skipinsert"` + Mode models.FilterMode `db:"mode"` + Name string `db:"name"` + FindFilter string `db:"find_filter"` + ObjectFilter string `db:"object_filter"` + UIOptions string `db:"ui_options"` +} + +func encodeJSONOrEmpty(v interface{}) string { + if v == nil { + return "" + } + + encoded, err := json.Marshal(v) + if err != nil { + logger.Errorf("error encoding json %v: %v", v, err) + } + + return string(encoded) +} + +func decodeJSON(s string, v interface{}) { + if s == "" { + return + } + + if err := json.Unmarshal([]byte(s), v); err != nil { + logger.Errorf("error decoding json %q: %v", s, err) + } +} + +func (r *savedFilterRow) fromSavedFilter(o models.SavedFilter) { + r.ID = o.ID + r.Mode = o.Mode + r.Name = o.Name + + // encode the filters as json + r.FindFilter = encodeJSONOrEmpty(o.FindFilter) + r.ObjectFilter = encodeJSONOrEmpty(o.ObjectFilter) + r.UIOptions = encodeJSONOrEmpty(o.UIOptions) +} + +func (r *savedFilterRow) resolve() *models.SavedFilter { + ret := &models.SavedFilter{ + ID: r.ID, + Mode: r.Mode, + Name: r.Name, + } + + // decode the filters from json + if r.FindFilter != "" { + ret.FindFilter = &models.FindFilterType{} + decodeJSON(r.FindFilter, &ret.FindFilter) + } + if r.ObjectFilter != "" { + ret.ObjectFilter = make(map[string]interface{}) + decodeJSON(r.ObjectFilter, &ret.ObjectFilter) + } + if r.UIOptions != "" { + ret.UIOptions = make(map[string]interface{}) + decodeJSON(r.UIOptions, &ret.UIOptions) + } + + return ret +} + +type SavedFilterStore struct { + repository + tableMgr *table +} + +func NewSavedFilterStore() *SavedFilterStore { + return &SavedFilterStore{ + repository: repository{ + tableName: savedFilterTable, + idColumn: idColumn, + }, + tableMgr: savedFilterTableMgr, + } +} + +func (qb *SavedFilterStore) table() exp.IdentifierExpression { + return qb.tableMgr.table +} + +func (qb *SavedFilterStore) selectDataset() *goqu.SelectDataset { + return dialect.From(qb.table()).Select(qb.table().All()) +} + +func (qb *SavedFilterStore) Create(ctx context.Context, newObject *models.SavedFilter) error { + var r savedFilterRow + r.fromSavedFilter(*newObject) + + id, err := qb.tableMgr.insertID(ctx, r) + if err != nil { + return err + } + + updated, err := qb.Find(ctx, id) + if err != nil { + return fmt.Errorf("finding after create: %w", err) + } + + *newObject = *updated + + return nil +} + +func (qb *SavedFilterStore) Update(ctx context.Context, updatedObject *models.SavedFilter) error { + var r savedFilterRow + r.fromSavedFilter(*updatedObject) + + if err := qb.tableMgr.updateByID(ctx, updatedObject.ID, r); err != nil { + return err + } + + return nil +} + +func (qb *SavedFilterStore) Destroy(ctx context.Context, id int) error { + return qb.destroyExisting(ctx, []int{id}) +} + +// returns nil, nil if not found +func (qb *SavedFilterStore) Find(ctx context.Context, id int) (*models.SavedFilter, error) { + ret, err := qb.find(ctx, id) + if errors.Is(err, sql.ErrNoRows) { + return nil, nil + } + return ret, err +} + +func (qb *SavedFilterStore) FindMany(ctx context.Context, ids []int, ignoreNotFound bool) ([]*models.SavedFilter, error) { + ret := make([]*models.SavedFilter, len(ids)) + + if len(ids) == 0 { + return ret, nil + } + + table := qb.table() + q := qb.selectDataset().Prepared(true).Where(table.Col(idColumn).In(ids)) + unsorted, err := qb.getMany(ctx, q) + if err != nil { + return nil, err + } + + for _, s := range unsorted { + i := slices.Index(ids, s.ID) + ret[i] = s + } + + if !ignoreNotFound { + for i := range ret { + if ret[i] == nil { + return nil, fmt.Errorf("filter with id %d not found", ids[i]) + } + } + } + + return ret, nil +} + +// returns nil, sql.ErrNoRows if not found +func (qb *SavedFilterStore) find(ctx context.Context, id int) (*models.SavedFilter, error) { + q := qb.selectDataset().Where(qb.tableMgr.byID(id)) + + ret, err := qb.get(ctx, q) + if err != nil { + return nil, err + } + + return ret, nil +} + +func (qb *SavedFilterStore) get(ctx context.Context, q *goqu.SelectDataset) (*models.SavedFilter, error) { + ret, err := qb.getMany(ctx, q) + if err != nil { + return nil, err + } + + if len(ret) == 0 { + return nil, sql.ErrNoRows + } + + return ret[0], nil +} + +func (qb *SavedFilterStore) getMany(ctx context.Context, q *goqu.SelectDataset) ([]*models.SavedFilter, error) { + const single = false + var ret []*models.SavedFilter + if err := queryFunc(ctx, q, single, func(r *sqlx.Rows) error { + var f savedFilterRow + if err := r.StructScan(&f); err != nil { + return err + } + + s := f.resolve() + + ret = append(ret, s) + return nil + }); err != nil { + return nil, err + } + + return ret, nil +} + +func (qb *SavedFilterStore) FindByMode(ctx context.Context, mode models.FilterMode) ([]*models.SavedFilter, error) { + // SELECT * FROM %s WHERE mode = ? AND name != ? ORDER BY name ASC + table := qb.table() + + // TODO - querying on groups needs to include movies + // remove this when we migrate to remove the movies filter mode in the database + var whereClause exp.Expression + + if mode == models.FilterModeGroups || mode == models.FilterModeMovies { + whereClause = goqu.Or( + table.Col("mode").Eq(models.FilterModeGroups), + table.Col("mode").Eq(models.FilterModeMovies), + ) + } else { + whereClause = table.Col("mode").Eq(mode) + } + + sq := qb.selectDataset().Prepared(true).Where(whereClause).Order(table.Col("name").Asc()) + ret, err := qb.getMany(ctx, sq) + + if err != nil { + return nil, err + } + + return ret, nil +} + +func (qb *SavedFilterStore) All(ctx context.Context) ([]*models.SavedFilter, error) { + return qb.getMany(ctx, qb.selectDataset()) +} diff --git a/pkg/postgres/savepoint.go b/pkg/postgres/savepoint.go new file mode 100644 index 0000000000..5d23fced6b --- /dev/null +++ b/pkg/postgres/savepoint.go @@ -0,0 +1,54 @@ +package postgres + +import ( + "context" + "fmt" + + "github.com/stashapp/stash/pkg/hash" +) + +type SavepointAction func(ctx context.Context) error + +const savePointPrefix = "savepoint_" // prefix for savepoint + +// Encapsulates an action in a savepoint +// Its mostly used to rollback if an error occurred in postgres, as errors in postgres cancel the transaction. +func withSavepoint(ctx context.Context, action SavepointAction) error { + tx, err := getTx(ctx) + if err != nil { + return err + } + + // Generate savepoint + rnd, err := hash.GenerateRandomKey(64) + if err != nil { + return err + } + + // Sqlite needs some letters infront of the identifier + rnd = savePointPrefix + rnd + + // Create a savepoint + _, err = tx.Exec("SAVEPOINT " + rnd) + if err != nil { + return fmt.Errorf("failed to create savepoint: %w", err) + } + + // Execute the action + err = action(ctx) + if err != nil { + // Rollback to savepoint on error + if _, rbErr := tx.Exec("ROLLBACK TO SAVEPOINT " + rnd); rbErr != nil { + return fmt.Errorf("action failed and rollback to savepoint failed: %w", rbErr) + } + return fmt.Errorf("action failed: %w", err) + } + + // Release the savepoint on success + _, err = tx.Exec("RELEASE SAVEPOINT " + rnd) + if err != nil { + return fmt.Errorf("failed to release savepoint: %w", err) + } + + return nil +} diff --git a/pkg/postgres/scene.go b/pkg/postgres/scene.go new file mode 100644 index 0000000000..95c32d4021 --- /dev/null +++ b/pkg/postgres/scene.go @@ -0,0 +1,1528 @@ +package postgres + +import ( + "context" + "database/sql" + "errors" + "fmt" + "path/filepath" + "slices" + "sort" + "strconv" + "strings" + + "github.com/doug-martin/goqu/v9" + "github.com/doug-martin/goqu/v9/exp" + "github.com/jmoiron/sqlx" + "gopkg.in/guregu/null.v4" + "gopkg.in/guregu/null.v4/zero" + + "github.com/stashapp/stash/pkg/models" + "github.com/stashapp/stash/pkg/sliceutil" + "github.com/stashapp/stash/pkg/utils" +) + +const ( + sceneTable = "scenes" + scenesFilesTable = "scenes_files" + sceneIDColumn = "scene_id" + performersScenesTable = "performers_scenes" + scenesTagsTable = "scenes_tags" + scenesGalleriesTable = "scenes_galleries" + groupsScenesTable = "groups_scenes" + scenesURLsTable = "scene_urls" + sceneURLColumn = "url" + scenesViewDatesTable = "scenes_view_dates" + sceneViewDateColumn = "view_date" + scenesODatesTable = "scenes_o_dates" + sceneODateColumn = "o_date" + + sceneCoverBlobColumn = "cover_blob" +) + +var findExactDuplicateQuery = ` +SELECT STRING_AGG(DISTINCT scene_id::TEXT, ',') as ids +FROM ( + SELECT scenes.id as scene_id + , video_files.duration as file_duration + , files.size as file_size + , files_fingerprints.fingerprint as phash + , abs(max(video_files.duration) OVER (PARTITION by files_fingerprints.fingerprint) - video_files.duration) as durationDiff + FROM scenes + INNER JOIN scenes_files ON (scenes.id = scenes_files.scene_id) + INNER JOIN files ON (scenes_files.file_id = files.id) + INNER JOIN files_fingerprints ON (scenes_files.file_id = files_fingerprints.file_id AND files_fingerprints.type = 'phash') + INNER JOIN video_files ON (files.id = video_files.file_id) +) as subq +WHERE durationDiff <= $1 + OR $1 < 0 -- Always TRUE if the parameter is negative. + -- That will disable the durationDiff checking. +GROUP BY phash +HAVING COUNT(phash) > 1 + AND COUNT(DISTINCT scene_id) > 1 +ORDER BY SUM(file_size) DESC; +` + +var findAllPhashesQuery = ` +SELECT scenes.id as id + , files_fingerprints.fingerprint as phash + , video_files.duration as duration +FROM scenes +INNER JOIN scenes_files ON (scenes.id = scenes_files.scene_id) +INNER JOIN files ON (scenes_files.file_id = files.id) +INNER JOIN files_fingerprints ON (scenes_files.file_id = files_fingerprints.file_id AND files_fingerprints.type = 'phash') +INNER JOIN video_files ON (files.id = video_files.file_id) +ORDER BY files.size DESC; +` + +type sceneRow struct { + ID int `db:"id" goqu:"skipinsert"` + Title zero.String `db:"title"` + Code zero.String `db:"code"` + Details zero.String `db:"details"` + Director zero.String `db:"director"` + Date NullDate `db:"date"` + DatePrecision null.Int `db:"date_precision"` + // expressed as 1-100 + Rating null.Int `db:"rating"` + Organized bool `db:"organized"` + StudioID null.Int `db:"studio_id,omitempty"` + CreatedAt Timestamp `db:"created_at"` + UpdatedAt Timestamp `db:"updated_at"` + ResumeTime float64 `db:"resume_time"` + PlayDuration float64 `db:"play_duration"` + + // not used in resolutions or updates + CoverBlob zero.String `db:"cover_blob"` +} + +func (r *sceneRow) fromScene(o models.Scene) { + r.ID = o.ID + r.Title = zero.StringFrom(o.Title) + r.Code = zero.StringFrom(o.Code) + r.Details = zero.StringFrom(o.Details) + r.Director = zero.StringFrom(o.Director) + r.Date = NullDateFromDatePtr(o.Date) + r.DatePrecision = datePrecisionFromDatePtr(o.Date) + r.Rating = intFromPtr(o.Rating) + r.Organized = o.Organized + r.StudioID = intFromPtr(o.StudioID) + r.CreatedAt = Timestamp{Timestamp: o.CreatedAt} + r.UpdatedAt = Timestamp{Timestamp: o.UpdatedAt} + r.ResumeTime = o.ResumeTime + r.PlayDuration = o.PlayDuration +} + +type sceneQueryRow struct { + sceneRow + PrimaryFileID null.Int `db:"primary_file_id"` + PrimaryFileFolderPath zero.String `db:"primary_file_folder_path"` + PrimaryFileBasename zero.String `db:"primary_file_basename"` + PrimaryFileOshash zero.String `db:"primary_file_oshash"` + PrimaryFileChecksum zero.String `db:"primary_file_checksum"` +} + +func (r *sceneQueryRow) resolve() *models.Scene { + ret := &models.Scene{ + ID: r.ID, + Title: r.Title.String, + Code: r.Code.String, + Details: r.Details.String, + Director: r.Director.String, + Date: r.Date.DatePtr(r.DatePrecision), + Rating: nullIntPtr(r.Rating), + Organized: r.Organized, + StudioID: nullIntPtr(r.StudioID), + + PrimaryFileID: nullIntFileIDPtr(r.PrimaryFileID), + OSHash: r.PrimaryFileOshash.String, + Checksum: r.PrimaryFileChecksum.String, + + CreatedAt: r.CreatedAt.Timestamp.UTC(), + UpdatedAt: r.UpdatedAt.Timestamp.UTC(), + + ResumeTime: r.ResumeTime, + PlayDuration: r.PlayDuration, + } + + if r.PrimaryFileFolderPath.Valid && r.PrimaryFileBasename.Valid { + ret.Path = filepath.Join(r.PrimaryFileFolderPath.String, r.PrimaryFileBasename.String) + } + + return ret +} + +type sceneRowRecord struct { + updateRecord +} + +func (r *sceneRowRecord) fromPartial(o models.ScenePartial) { + r.setNullString("title", o.Title) + r.setNullString("code", o.Code) + r.setNullString("details", o.Details) + r.setNullString("director", o.Director) + r.setNullDate("date", "date_precision", o.Date) + r.setNullInt("rating", o.Rating) + r.setBool("organized", o.Organized) + r.setNullInt("studio_id", o.StudioID) + r.setTimestamp("created_at", o.CreatedAt) + r.setTimestamp("updated_at", o.UpdatedAt) + r.setFloat64("resume_time", o.ResumeTime) + r.setFloat64("play_duration", o.PlayDuration) +} + +type sceneRepositoryType struct { + repository + galleries joinRepository + tags joinRepository + performers joinRepository + groups repository + + files filesRepository + + stashIDs stashIDRepository +} + +var ( + sceneRepository = sceneRepositoryType{ + repository: repository{ + tableName: sceneTable, + idColumn: idColumn, + }, + galleries: joinRepository{ + repository: repository{ + tableName: scenesGalleriesTable, + idColumn: sceneIDColumn, + }, + fkColumn: galleryIDColumn, + }, + tags: joinRepository{ + repository: repository{ + tableName: scenesTagsTable, + idColumn: sceneIDColumn, + }, + fkColumn: tagIDColumn, + foreignTable: tagTable, + orderBy: tagTableSortSQL, + }, + performers: joinRepository{ + repository: repository{ + tableName: performersScenesTable, + idColumn: sceneIDColumn, + }, + fkColumn: performerIDColumn, + }, + groups: repository{ + tableName: groupsScenesTable, + idColumn: sceneIDColumn, + }, + files: filesRepository{ + repository: repository{ + tableName: scenesFilesTable, + idColumn: sceneIDColumn, + }, + }, + stashIDs: stashIDRepository{ + repository{ + tableName: "scene_stash_ids", + idColumn: sceneIDColumn, + }, + }, + } +) + +type SceneStore struct { + blobJoinQueryBuilder + + tableMgr *table + oDateManager + viewDateManager + + repo *storeRepository +} + +func NewSceneStore(r *storeRepository, blobStore *BlobStore) *SceneStore { + return &SceneStore{ + blobJoinQueryBuilder: blobJoinQueryBuilder{ + blobStore: blobStore, + joinTable: sceneTable, + }, + + tableMgr: sceneTableMgr, + viewDateManager: viewDateManager{scenesViewTableMgr}, + oDateManager: oDateManager{scenesOTableMgr}, + repo: r, + } +} + +func (qb *SceneStore) table() exp.IdentifierExpression { + return qb.tableMgr.table +} + +func (qb *SceneStore) selectDataset() *goqu.SelectDataset { + table := qb.table() + files := fileTableMgr.table + folders := folderTableMgr.table + checksum := fingerprintTableMgr.table.As("fingerprint_md5") + oshash := fingerprintTableMgr.table.As("fingerprint_oshash") + + return dialect.From(table).LeftJoin( + scenesFilesJoinTable, + goqu.On( + scenesFilesJoinTable.Col(sceneIDColumn).Eq(table.Col(idColumn)), + scenesFilesJoinTable.Col("primary").IsTrue(), + ), + ).LeftJoin( + files, + goqu.On(files.Col(idColumn).Eq(scenesFilesJoinTable.Col(fileIDColumn))), + ).LeftJoin( + folders, + goqu.On(folders.Col(idColumn).Eq(files.Col("parent_folder_id"))), + ).LeftJoin( + checksum, + goqu.On( + checksum.Col(fileIDColumn).Eq(scenesFilesJoinTable.Col(fileIDColumn)), + checksum.Col("type").Eq(models.FingerprintTypeMD5), + ), + ).LeftJoin( + oshash, + goqu.On( + oshash.Col(fileIDColumn).Eq(scenesFilesJoinTable.Col(fileIDColumn)), + oshash.Col("type").Eq(models.FingerprintTypeOshash), + ), + ).Select( + qb.table().All(), + scenesFilesJoinTable.Col(fileIDColumn).As("primary_file_id"), + folders.Col("path").As("primary_file_folder_path"), + files.Col("basename").As("primary_file_basename"), + checksum.Col("fingerprint").As("primary_file_checksum"), + oshash.Col("fingerprint").As("primary_file_oshash"), + ) +} + +func (qb *SceneStore) Create(ctx context.Context, newObject *models.Scene, fileIDs []models.FileID) error { + var r sceneRow + r.fromScene(*newObject) + + id, err := qb.tableMgr.insertID(ctx, r) + if err != nil { + return err + } + + if len(fileIDs) > 0 { + const firstPrimary = true + if err := scenesFilesTableMgr.insertJoins(ctx, id, firstPrimary, fileIDs); err != nil { + return err + } + } + + if newObject.URLs.Loaded() { + const startPos = 0 + if err := scenesURLsTableMgr.insertJoins(ctx, id, startPos, newObject.URLs.List()); err != nil { + return err + } + } + + if newObject.PerformerIDs.Loaded() { + if err := scenesPerformersTableMgr.insertJoins(ctx, id, newObject.PerformerIDs.List()); err != nil { + return err + } + } + if newObject.TagIDs.Loaded() { + if err := scenesTagsTableMgr.insertJoins(ctx, id, newObject.TagIDs.List()); err != nil { + return err + } + } + + if newObject.GalleryIDs.Loaded() { + if err := scenesGalleriesTableMgr.insertJoins(ctx, id, newObject.GalleryIDs.List()); err != nil { + return err + } + } + + if newObject.StashIDs.Loaded() { + if err := scenesStashIDsTableMgr.insertJoins(ctx, id, newObject.StashIDs.List()); err != nil { + return err + } + } + + if newObject.Groups.Loaded() { + if err := scenesGroupsTableMgr.insertJoins(ctx, id, newObject.Groups.List()); err != nil { + return err + } + } + + updated, err := qb.find(ctx, id) + if err != nil { + return fmt.Errorf("finding after create: %w", err) + } + + *newObject = *updated + + return nil +} + +func (qb *SceneStore) UpdatePartial(ctx context.Context, id int, partial models.ScenePartial) (*models.Scene, error) { + r := sceneRowRecord{ + updateRecord{ + Record: make(exp.Record), + }, + } + + r.fromPartial(partial) + + if len(r.Record) > 0 { + if err := qb.tableMgr.updateByID(ctx, id, r.Record); err != nil { + return nil, err + } + } + + if partial.URLs != nil { + if err := scenesURLsTableMgr.modifyJoins(ctx, id, partial.URLs.Values, partial.URLs.Mode); err != nil { + return nil, err + } + } + if partial.PerformerIDs != nil { + if err := scenesPerformersTableMgr.modifyJoins(ctx, id, partial.PerformerIDs.IDs, partial.PerformerIDs.Mode); err != nil { + return nil, err + } + } + if partial.TagIDs != nil { + if err := scenesTagsTableMgr.modifyJoins(ctx, id, partial.TagIDs.IDs, partial.TagIDs.Mode); err != nil { + return nil, err + } + } + if partial.GalleryIDs != nil { + if err := scenesGalleriesTableMgr.modifyJoins(ctx, id, partial.GalleryIDs.IDs, partial.GalleryIDs.Mode); err != nil { + return nil, err + } + } + if partial.StashIDs != nil { + if err := scenesStashIDsTableMgr.modifyJoins(ctx, id, partial.StashIDs.StashIDs, partial.StashIDs.Mode); err != nil { + return nil, err + } + } + if partial.GroupIDs != nil { + if err := scenesGroupsTableMgr.modifyJoins(ctx, id, partial.GroupIDs.Groups, partial.GroupIDs.Mode); err != nil { + return nil, err + } + } + if partial.PrimaryFileID != nil { + if err := scenesFilesTableMgr.setPrimary(ctx, id, *partial.PrimaryFileID); err != nil { + return nil, err + } + } + + return qb.find(ctx, id) +} + +func (qb *SceneStore) Update(ctx context.Context, updatedObject *models.Scene) error { + var r sceneRow + r.fromScene(*updatedObject) + + if err := qb.tableMgr.updateByID(ctx, updatedObject.ID, r); err != nil { + return err + } + + if updatedObject.URLs.Loaded() { + if err := scenesURLsTableMgr.replaceJoins(ctx, updatedObject.ID, updatedObject.URLs.List()); err != nil { + return err + } + } + + if updatedObject.PerformerIDs.Loaded() { + if err := scenesPerformersTableMgr.replaceJoins(ctx, updatedObject.ID, updatedObject.PerformerIDs.List()); err != nil { + return err + } + } + + if updatedObject.TagIDs.Loaded() { + if err := scenesTagsTableMgr.replaceJoins(ctx, updatedObject.ID, updatedObject.TagIDs.List()); err != nil { + return err + } + } + + if updatedObject.GalleryIDs.Loaded() { + if err := scenesGalleriesTableMgr.replaceJoins(ctx, updatedObject.ID, updatedObject.GalleryIDs.List()); err != nil { + return err + } + } + + if updatedObject.StashIDs.Loaded() { + if err := scenesStashIDsTableMgr.replaceJoins(ctx, updatedObject.ID, updatedObject.StashIDs.List()); err != nil { + return err + } + } + + if updatedObject.Groups.Loaded() { + if err := scenesGroupsTableMgr.replaceJoins(ctx, updatedObject.ID, updatedObject.Groups.List()); err != nil { + return err + } + } + + if updatedObject.Files.Loaded() { + fileIDs := make([]models.FileID, len(updatedObject.Files.List())) + for i, f := range updatedObject.Files.List() { + fileIDs[i] = f.ID + } + + if err := scenesFilesTableMgr.replaceJoins(ctx, updatedObject.ID, fileIDs); err != nil { + return err + } + } + + return nil +} + +func (qb *SceneStore) Destroy(ctx context.Context, id int) error { + // must handle image checksums manually + if err := qb.destroyCover(ctx, id); err != nil { + return err + } + + // scene markers should be handled prior to calling destroy + // galleries should be handled prior to calling destroy + + return qb.tableMgr.destroyExisting(ctx, []int{id}) +} + +// returns nil, nil if not found +func (qb *SceneStore) Find(ctx context.Context, id int) (*models.Scene, error) { + ret, err := qb.find(ctx, id) + if errors.Is(err, sql.ErrNoRows) { + return nil, nil + } + return ret, err +} + +// FindByIDs finds multiple scenes by their IDs. +// No check is made to see if the scenes exist, and the order of the returned scenes +// is not guaranteed to be the same as the order of the input IDs. +func (qb *SceneStore) FindByIDs(ctx context.Context, ids []int) ([]*models.Scene, error) { + scenes := make([]*models.Scene, 0, len(ids)) + + if len(ids) == 0 { + return scenes, nil + } + + table := qb.table() + if err := batchExec(ids, defaultBatchSize, func(batch []int) error { + q := qb.selectDataset().Prepared(true).Where(table.Col(idColumn).In(batch)) + unsorted, err := qb.getMany(ctx, q) + if err != nil { + return err + } + + scenes = append(scenes, unsorted...) + + return nil + }); err != nil { + return nil, err + } + + return scenes, nil +} + +func (qb *SceneStore) FindMany(ctx context.Context, ids []int) ([]*models.Scene, error) { + scenes := make([]*models.Scene, len(ids)) + + if len(ids) == 0 { + return scenes, nil + } + + unsorted, err := qb.FindByIDs(ctx, ids) + if err != nil { + return nil, err + } + + for _, s := range unsorted { + i := slices.Index(ids, s.ID) + scenes[i] = s + } + + for i := range scenes { + if scenes[i] == nil { + return nil, fmt.Errorf("scene with id %d not found", ids[i]) + } + } + + return scenes, nil +} + +// returns nil, sql.ErrNoRows if not found +func (qb *SceneStore) find(ctx context.Context, id int) (*models.Scene, error) { + q := qb.selectDataset().Where(qb.tableMgr.byID(id)) + + ret, err := qb.get(ctx, q) + if err != nil { + return nil, err + } + + return ret, nil +} + +func (qb *SceneStore) findBySubquery(ctx context.Context, sq *goqu.SelectDataset) ([]*models.Scene, error) { + table := qb.table() + + q := qb.selectDataset().Where( + table.Col(idColumn).Eq( + sq, + ), + ) + + return qb.getMany(ctx, q) +} + +// returns nil, sql.ErrNoRows if not found +func (qb *SceneStore) get(ctx context.Context, q *goqu.SelectDataset) (*models.Scene, error) { + ret, err := qb.getMany(ctx, q) + if err != nil { + return nil, err + } + + if len(ret) == 0 { + return nil, sql.ErrNoRows + } + + return ret[0], nil +} + +func (qb *SceneStore) getMany(ctx context.Context, q *goqu.SelectDataset) ([]*models.Scene, error) { + const single = false + var ret []*models.Scene + var lastID int + if err := queryFunc(ctx, q, single, func(r *sqlx.Rows) error { + var f sceneQueryRow + if err := r.StructScan(&f); err != nil { + return err + } + + s := f.resolve() + if s.ID == lastID { + return fmt.Errorf("internal error: multiple rows returned for single scene id %d", s.ID) + } + lastID = s.ID + + ret = append(ret, s) + return nil + }); err != nil { + return nil, err + } + + return ret, nil +} + +func (qb *SceneStore) GetFiles(ctx context.Context, id int) ([]*models.VideoFile, error) { + fileIDs, err := sceneRepository.files.get(ctx, id) + if err != nil { + return nil, err + } + + // use fileStore to load files + files, err := qb.repo.File.Find(ctx, fileIDs...) + if err != nil { + return nil, err + } + + ret := make([]*models.VideoFile, len(files)) + for i, f := range files { + var ok bool + ret[i], ok = f.(*models.VideoFile) + if !ok { + return nil, fmt.Errorf("expected file to be *file.VideoFile not %T", f) + } + } + + return ret, nil +} + +func (qb *SceneStore) GetManyFileIDs(ctx context.Context, ids []int) ([][]models.FileID, error) { + const primaryOnly = false + return sceneRepository.files.getMany(ctx, ids, primaryOnly) +} + +func (qb *SceneStore) FindByFileID(ctx context.Context, fileID models.FileID) ([]*models.Scene, error) { + sq := dialect.From(scenesFilesJoinTable).Select(scenesFilesJoinTable.Col(sceneIDColumn)).Where( + scenesFilesJoinTable.Col(fileIDColumn).Eq(fileID), + ) + + ret, err := qb.findBySubquery(ctx, sq) + if err != nil { + return nil, fmt.Errorf("getting scenes by file id %d: %w", fileID, err) + } + + return ret, nil +} + +func (qb *SceneStore) FindByPrimaryFileID(ctx context.Context, fileID models.FileID) ([]*models.Scene, error) { + sq := dialect.From(scenesFilesJoinTable).Select(scenesFilesJoinTable.Col(sceneIDColumn)).Where( + scenesFilesJoinTable.Col(fileIDColumn).Eq(fileID), + scenesFilesJoinTable.Col("primary").IsTrue(), + ) + + ret, err := qb.findBySubquery(ctx, sq) + if err != nil { + return nil, fmt.Errorf("getting scenes by primary file id %d: %w", fileID, err) + } + + return ret, nil +} + +func (qb *SceneStore) CountByFileID(ctx context.Context, fileID models.FileID) (int, error) { + joinTable := scenesFilesJoinTable + + q := dialect.Select(goqu.COUNT("*")).From(joinTable).Where(joinTable.Col(fileIDColumn).Eq(fileID)) + return count(ctx, q) +} + +func (qb *SceneStore) FindByFingerprints(ctx context.Context, fp []models.Fingerprint) ([]*models.Scene, error) { + fingerprintTable := fingerprintTableMgr.table + + var ex []exp.Expression + + for _, v := range fp { + ex = append(ex, goqu.And( + fingerprintTable.Col("type").Eq(v.Type), + fingerprintTable.Col("fingerprint").Eq(v.Fingerprint), + )) + } + + sq := dialect.From(scenesFilesJoinTable). + InnerJoin( + fingerprintTable, + goqu.On(fingerprintTable.Col(fileIDColumn).Eq(scenesFilesJoinTable.Col(fileIDColumn))), + ). + Select(scenesFilesJoinTable.Col(sceneIDColumn)).Where(goqu.Or(ex...)) + + ret, err := qb.findBySubquery(ctx, sq) + if err != nil { + return nil, fmt.Errorf("getting scenes by fingerprints: %w", err) + } + + return ret, nil +} + +func (qb *SceneStore) FindByChecksum(ctx context.Context, checksum string) ([]*models.Scene, error) { + return qb.FindByFingerprints(ctx, []models.Fingerprint{ + { + Type: models.FingerprintTypeMD5, + Fingerprint: checksum, + }, + }) +} + +func (qb *SceneStore) FindByOSHash(ctx context.Context, oshash string) ([]*models.Scene, error) { + return qb.FindByFingerprints(ctx, []models.Fingerprint{ + { + Type: models.FingerprintTypeOshash, + Fingerprint: oshash, + }, + }) +} + +func (qb *SceneStore) FindByPath(ctx context.Context, p string) ([]*models.Scene, error) { + filesTable := fileTableMgr.table + foldersTable := folderTableMgr.table + basename := filepath.Base(p) + dir := filepath.Dir(p) + + // replace wildcards + basename = strings.ReplaceAll(basename, "*", "%") + dir = strings.ReplaceAll(dir, "*", "%") + + sq := dialect.From(scenesFilesJoinTable).InnerJoin( + filesTable, + goqu.On(filesTable.Col(idColumn).Eq(scenesFilesJoinTable.Col(fileIDColumn))), + ).InnerJoin( + foldersTable, + goqu.On(foldersTable.Col(idColumn).Eq(filesTable.Col("parent_folder_id"))), + ).Select(scenesFilesJoinTable.Col(sceneIDColumn)).Where( + foldersTable.Col("path").Like(dir), + filesTable.Col("basename").Like(basename), + ) + + ret, err := qb.findBySubquery(ctx, sq) + if err != nil && !errors.Is(err, sql.ErrNoRows) { + return nil, fmt.Errorf("getting scene by path %s: %w", p, err) + } + + return ret, nil +} + +func (qb *SceneStore) FindByPerformerID(ctx context.Context, performerID int) ([]*models.Scene, error) { + sq := dialect.From(scenesPerformersJoinTable).Select(scenesPerformersJoinTable.Col(sceneIDColumn)).Where( + scenesPerformersJoinTable.Col(performerIDColumn).Eq(performerID), + ) + ret, err := qb.findBySubquery(ctx, sq) + + if err != nil { + return nil, fmt.Errorf("getting scenes for performer %d: %w", performerID, err) + } + + return ret, nil +} + +func (qb *SceneStore) FindByGalleryID(ctx context.Context, galleryID int) ([]*models.Scene, error) { + sq := dialect.From(galleriesScenesJoinTable).Select(galleriesScenesJoinTable.Col(sceneIDColumn)).Where( + galleriesScenesJoinTable.Col(galleryIDColumn).Eq(galleryID), + ) + ret, err := qb.findBySubquery(ctx, sq) + + if err != nil { + return nil, fmt.Errorf("getting scenes for gallery %d: %w", galleryID, err) + } + + return ret, nil +} + +func (qb *SceneStore) CountByPerformerID(ctx context.Context, performerID int) (int, error) { + joinTable := scenesPerformersJoinTable + + q := dialect.Select(goqu.COUNT("*")).From(joinTable).Where(joinTable.Col(performerIDColumn).Eq(performerID)) + return count(ctx, q) +} + +func (qb *SceneStore) OCountByPerformerID(ctx context.Context, performerID int) (int, error) { + table := qb.table() + joinTable := scenesPerformersJoinTable + oHistoryTable := goqu.T(scenesODatesTable) + + q := dialect.Select(goqu.COUNT("*")).From(table).InnerJoin( + oHistoryTable, + goqu.On(table.Col(idColumn).Eq(oHistoryTable.Col(sceneIDColumn))), + ).InnerJoin( + joinTable, + goqu.On( + table.Col(idColumn).Eq(joinTable.Col(sceneIDColumn)), + ), + ).Where(joinTable.Col(performerIDColumn).Eq(performerID)) + + var ret int + if err := querySimple(ctx, q, &ret); err != nil { + return 0, err + } + + return ret, nil +} + +func (qb *SceneStore) OCountByGroupID(ctx context.Context, groupID int) (int, error) { + table := qb.table() + joinTable := scenesGroupsJoinTable + oHistoryTable := goqu.T(scenesODatesTable) + + q := dialect.Select(goqu.COUNT("*")).From(table).InnerJoin( + oHistoryTable, + goqu.On(table.Col(idColumn).Eq(oHistoryTable.Col(sceneIDColumn))), + ).InnerJoin( + joinTable, + goqu.On( + table.Col(idColumn).Eq(joinTable.Col(sceneIDColumn)), + ), + ).Where(joinTable.Col(groupIDColumn).Eq(groupID)) + + var ret int + if err := querySimple(ctx, q, &ret); err != nil { + return 0, err + } + + return ret, nil +} + +func (qb *SceneStore) OCountByStudioID(ctx context.Context, studioID int) (int, error) { + table := qb.table() + oHistoryTable := goqu.T(scenesODatesTable) + + q := dialect.Select(goqu.COUNT("*")).From(table).InnerJoin( + oHistoryTable, + goqu.On(table.Col(idColumn).Eq(oHistoryTable.Col(sceneIDColumn))), + ).Where(table.Col(studioIDColumn).Eq(studioID)) + + var ret int + if err := querySimple(ctx, q, &ret); err != nil { + return 0, err + } + + return ret, nil +} + +func (qb *SceneStore) FindByGroupID(ctx context.Context, groupID int) ([]*models.Scene, error) { + sq := dialect.From(scenesGroupsJoinTable).Select(scenesGroupsJoinTable.Col(sceneIDColumn)).Where( + scenesGroupsJoinTable.Col(groupIDColumn).Eq(groupID), + ) + ret, err := qb.findBySubquery(ctx, sq) + + if err != nil { + return nil, fmt.Errorf("getting scenes for group %d: %w", groupID, err) + } + + return ret, nil +} + +func (qb *SceneStore) Count(ctx context.Context) (int, error) { + q := dialect.Select(goqu.COUNT("*")).From(qb.table()) + return count(ctx, q) +} + +func (qb *SceneStore) Size(ctx context.Context) (float64, error) { + table := qb.table() + fileTable := fileTableMgr.table + q := dialect.Select( + goqu.COALESCE(goqu.SUM(fileTableMgr.table.Col("size")), 0), + ).From(table).InnerJoin( + scenesFilesJoinTable, + goqu.On(table.Col(idColumn).Eq(scenesFilesJoinTable.Col(sceneIDColumn))), + ).InnerJoin( + fileTable, + goqu.On(scenesFilesJoinTable.Col(fileIDColumn).Eq(fileTable.Col(idColumn))), + ) + var ret float64 + if err := querySimple(ctx, q, &ret); err != nil { + return 0, err + } + + return ret, nil +} + +func (qb *SceneStore) Duration(ctx context.Context) (float64, error) { + table := qb.table() + videoFileTable := videoFileTableMgr.table + + q := dialect.Select( + goqu.COALESCE(goqu.SUM(videoFileTable.Col("duration")), 0), + ).From(table).InnerJoin( + scenesFilesJoinTable, + goqu.On(scenesFilesJoinTable.Col("scene_id").Eq(table.Col(idColumn))), + ).InnerJoin( + videoFileTable, + goqu.On(videoFileTable.Col("file_id").Eq(scenesFilesJoinTable.Col("file_id"))), + ) + + var ret float64 + if err := querySimple(ctx, q, &ret); err != nil { + return 0, err + } + + return ret, nil +} + +func (qb *SceneStore) PlayDuration(ctx context.Context) (float64, error) { + table := qb.table() + + q := dialect.Select(goqu.COALESCE(goqu.SUM("play_duration"), 0)).From(table) + + var ret float64 + if err := querySimple(ctx, q, &ret); err != nil { + return 0, err + } + + return ret, nil +} + +// TODO - currently only used by unit test +func (qb *SceneStore) CountByStudioID(ctx context.Context, studioID int) (int, error) { + table := qb.table() + + q := dialect.Select(goqu.COUNT("*")).From(table).Where(table.Col(studioIDColumn).Eq(studioID)) + return count(ctx, q) +} + +func (qb *SceneStore) countMissingFingerprints(ctx context.Context, fpType string) (int, error) { + fpTable := fingerprintTableMgr.table.As("fingerprints_temp") + + q := dialect.From(scenesFilesJoinTable).LeftJoin( + fpTable, + goqu.On( + scenesFilesJoinTable.Col(fileIDColumn).Eq(fpTable.Col(fileIDColumn)), + fpTable.Col("type").Eq(fpType), + ), + ).Select(goqu.COUNT(goqu.DISTINCT(scenesFilesJoinTable.Col(sceneIDColumn)))).Where(fpTable.Col("fingerprint").IsNull()) + + return count(ctx, q) +} + +// CountMissingChecksum returns the number of scenes missing a checksum value. +func (qb *SceneStore) CountMissingChecksum(ctx context.Context) (int, error) { + return qb.countMissingFingerprints(ctx, "md5") +} + +// CountMissingOSHash returns the number of scenes missing an oshash value. +func (qb *SceneStore) CountMissingOSHash(ctx context.Context) (int, error) { + return qb.countMissingFingerprints(ctx, "oshash") +} + +func (qb *SceneStore) Wall(ctx context.Context, q *string) ([]*models.Scene, error) { + s := "" + if q != nil { + s = *q + } + + table := qb.table() + qq := qb.selectDataset().Prepared(true).Where(table.Col("details").ILike("%" + s + "%")).Order(goqu.L("RANDOM()").Asc()).Limit(80) + return qb.getMany(ctx, qq) +} + +func (qb *SceneStore) All(ctx context.Context) ([]*models.Scene, error) { + table := qb.table() + fileTable := fileTableMgr.table + folderTable := folderTableMgr.table + + return qb.getMany(ctx, qb.selectDataset().Order( + folderTable.Col("path").Asc(), + fileTable.Col("basename").Asc(), + table.Col("date").Asc(), + )) +} + +func (qb *SceneStore) makeQuery(ctx context.Context, sceneFilter *models.SceneFilterType, findFilter *models.FindFilterType) (*queryBuilder, error) { + if sceneFilter == nil { + sceneFilter = &models.SceneFilterType{} + } + if findFilter == nil { + findFilter = &models.FindFilterType{} + } + + query := sceneRepository.newQuery() + distinctIDs(&query, sceneTable) + + if q := findFilter.Q; q != nil && *q != "" { + query.addJoins( + join{ + table: scenesFilesTable, + onClause: "scenes_files.scene_id = scenes.id AND scenes_files.\"primary\" = true", + }, + join{ + table: fileTable, + onClause: "scenes_files.file_id = files.id", + }, + join{ + table: folderTable, + onClause: "files.parent_folder_id = folders.id", + }, + join{ + table: fingerprintTable, + onClause: "files_fingerprints.file_id = scenes_files.file_id", + }, + join{ + table: sceneMarkerTable, + onClause: "scene_markers.scene_id = scenes.id", + }, + ) + + filepathColumn := "folders.path || '" + string(filepath.Separator) + "' || files.basename" + searchColumns := []string{"scenes.title", "scenes.details", filepathColumn, "files_fingerprints.fingerprint", "scene_markers.title"} + query.parseQueryString(searchColumns, *q) + } + + filter := filterBuilderFromHandler(ctx, &sceneFilterHandler{ + sceneFilter: sceneFilter, + }) + + if err := query.addFilter(filter); err != nil { + return nil, err + } + + if err := qb.setSceneSort(&query, findFilter); err != nil { + return nil, err + } + query.pagination = getPagination(findFilter) + + return &query, nil +} + +func (qb *SceneStore) Query(ctx context.Context, options models.SceneQueryOptions) (*models.SceneQueryResult, error) { + query, err := qb.makeQuery(ctx, options.SceneFilter, options.FindFilter) + if err != nil { + return nil, err + } + + result, err := qb.queryGroupedFields(ctx, options, *query) + if err != nil { + return nil, fmt.Errorf("error querying aggregate fields: %w", err) + } + + return result, nil +} + +func (qb *SceneStore) queryGroupedFields(ctx context.Context, options models.SceneQueryOptions, query queryBuilder) (*models.SceneQueryResult, error) { + // Add necessary joins and columns for aggregation + if options.Count { + query.addColumn("COUNT(*) OVER () AS total_count") + } + if options.TotalDuration { + query.addJoins( + join{ + table: scenesFilesTable, + onClause: "scenes_files.scene_id = scenes.id AND scenes_files.\"primary\" = true", + }, + join{ + table: videoFileTable, + onClause: "scenes_files.file_id = video_files.file_id", + }, + ) + query.addGroupBy("video_files.duration") + query.addColumn("SUM(COALESCE(files.size, 0)) OVER () AS total_size") + } + if options.TotalSize { + query.addJoins( + join{ + table: scenesFilesTable, + onClause: "scenes_files.scene_id = scenes.id AND scenes_files.\"primary\" = true", + }, + join{ + table: fileTable, + onClause: "scenes_files.file_id = files.id", + }, + ) + query.addGroupBy("files.size") + query.addColumn("SUM(COALESCE(video_files.duration, 0)) OVER () AS total_custom") + } + + // Support counting only + const includeSortPagination = true + var countOnly = options.FindFilter != nil && options.FindFilter.IsCounting() + + // Execute aggregate query + var obj *RowsWithCounts + var err error + if obj, err = sceneRepository.runIdsWithCount(ctx, query.toSQL(includeSortPagination), query.args); err != nil { + return nil, err + } + + // Prepare result + ret := models.NewSceneQueryResult(qb) + if len(obj.IDs) == 0 { + return ret, nil + } + + if !countOnly { + ret.IDs = obj.IDs + } + if options.Count { + ret.Count = int(obj.TotalCount.Int64) + } + if options.TotalSize { + ret.TotalSize = obj.TotalSize.Float64 + } + if options.TotalDuration { + ret.TotalDuration = obj.TotalCustom.Float64 + } + return ret, nil +} + +func (qb *SceneStore) QueryCount(ctx context.Context, sceneFilter *models.SceneFilterType, findFilter *models.FindFilterType) (int, error) { + query, err := qb.makeQuery(ctx, sceneFilter, findFilter) + if err != nil { + return 0, err + } + + return query.executeCount(ctx) +} + +var sceneSortOptions = sortOptions{ + "bitrate", + "created_at", + "code", + "date", + "file_count", + "filesize", + "duration", + "file_mod_time", + "framerate", + "group_scene_number", + "id", + "interactive", + "interactive_speed", + "last_o_at", + "last_played_at", + "movie_scene_number", + "o_counter", + "organized", + "performer_count", + "play_count", + "play_duration", + "resume_time", + "path", + "perceptual_similarity", + "random", + "rating", + "studio", + "tag_count", + "title", + "updated_at", + "performer_age", +} + +func (qb *SceneStore) setSceneSort(query *queryBuilder, findFilter *models.FindFilterType) error { + if findFilter == nil || findFilter.Sort == nil || *findFilter.Sort == "" { + return nil + } + sort := findFilter.GetSort("title") + + // CVE-2024-32231 - ensure sort is in the list of allowed sorts + if err := sceneSortOptions.validateSort(sort); err != nil { + return err + } + + addFileTable := func() { + query.addJoins( + join{ + sort: true, + table: scenesFilesTable, + onClause: "scenes_files.scene_id = scenes.id AND scenes_files.\"primary\" = true", + }, + join{ + sort: true, + table: fileTable, + onClause: "scenes_files.file_id = files.id", + }, + ) + } + + addVideoFileTable := func() { + addFileTable() + query.addJoins( + join{ + sort: true, + table: videoFileTable, + onClause: "video_files.file_id = scenes_files.file_id", + }, + ) + } + + addFolderTable := func() { + query.addJoins( + join{ + sort: true, + table: folderTable, + onClause: "files.parent_folder_id = folders.id", + }, + ) + } + + direction := findFilter.GetDirection() + switch sort { + case "movie_scene_number": + query.joinSort(groupsScenesTable, "", "scenes.id = groups_scenes.scene_id") + add, group := getSort("scene_index", direction, groupsScenesTable) + query.sort += add + query.addGroupBy(group...) + case "group_scene_number": + query.joinSort(groupsScenesTable, "scene_group", "scenes.id = scene_group.scene_id") + add, group := getSort("scene_index", direction, "scene_group") + query.sort += add + query.addGroupBy(group...) + case "tag_count": + query.sort += getCountSort(sceneTable, scenesTagsTable, sceneIDColumn, direction) + case "performer_count": + query.sort += getCountSort(sceneTable, performersScenesTable, sceneIDColumn, direction) + case "file_count": + query.sort += getCountSort(sceneTable, scenesFilesTable, sceneIDColumn, direction) + case "path": + // special handling for path + addFileTable() + addFolderTable() + query.sort += fmt.Sprintf(" ORDER BY COALESCE(folders.path, '') || COALESCE(files.basename, '') COLLATE NATURAL_CI %s", direction) + query.addGroupBy("folders.path", "files.basename") + case "perceptual_similarity": + // special handling for phash + addFileTable() + query.addJoins( + join{ + sort: true, + table: fingerprintTable, + as: "fingerprints_phash", + onClause: "scenes_files.file_id = fingerprints_phash.file_id AND fingerprints_phash.type = 'phash'", + }, + ) + + query.sort += " ORDER BY fingerprints_phash.fingerprint " + direction + ", files.size DESC" + query.addGroupBy("fingerprints_phash.fingerprint", "files.size") + case "bitrate": + sort = "bit_rate" + addVideoFileTable() + add, group := getSort(sort, direction, videoFileTable) + query.sort += add + query.addGroupBy(group...) + case "file_mod_time": + sort = "mod_time" + addFileTable() + add, agg := getSort(sort, direction, fileTable) + query.sort += add + query.addGroupBy(agg...) + case "framerate": + sort = "frame_rate" + addVideoFileTable() + add, agg := getSort(sort, direction, videoFileTable) + query.sort += add + query.addGroupBy(agg...) + case "resolution": + addVideoFileTable() + query.sort += fmt.Sprintf(" ORDER BY MIN(%s.width, %s.height) %s", videoFileTable, videoFileTable, getSortDirection(direction)) + case "filesize": + addFileTable() + add, agg := getSort(sort, direction, fileTable) + query.sort += add + query.addGroupBy(agg...) + case "duration": + addVideoFileTable() + add, agg := getSort(sort, direction, videoFileTable) + query.sort += add + query.addGroupBy(agg...) + case "interactive", "interactive_speed": + addVideoFileTable() + add, agg := getSort(sort, direction, videoFileTable) + query.sort += add + query.addGroupBy(agg...) + case "title": + addFileTable() + addFolderTable() + query.sort += " ORDER BY COALESCE(scenes.title, files.basename) COLLATE NATURAL_CI " + direction + ", folders.path COLLATE NATURAL_CI " + direction + query.addGroupBy("scenes.title", "files.basename", "folders.path") + case "play_count": + query.sort += getCountSort(sceneTable, scenesViewDatesTable, sceneIDColumn, direction) + case "last_played_at": + query.sort += fmt.Sprintf(" ORDER BY (SELECT MAX(view_date) FROM %s AS sort WHERE sort.%s = %s.id) %s", scenesViewDatesTable, sceneIDColumn, sceneTable, getSortDirection(direction)) + case "last_o_at": + query.sort += fmt.Sprintf(" ORDER BY (SELECT MAX(o_date) FROM %s AS sort WHERE sort.%s = %s.id) %s", scenesODatesTable, sceneIDColumn, sceneTable, getSortDirection(direction)) + case "o_counter": + query.sort += getCountSort(sceneTable, scenesODatesTable, sceneIDColumn, direction) + case "performer_age": + // Looking at the youngest performer by default + aggregation := "MIN" + if direction == "DESC" { + // When sorting by performer_'s age DESC, I should consider the oldest performer instead + aggregation = "MAX" + } + fallback := "NULL" + if direction == "ASC" { + // When sorting ascending, NULLs are first by default. Coalescing to the MAX int value supported by sqlite + fallback = "9223372036854775807" + } + query.sort += fmt.Sprintf( + " ORDER BY (SELECT COALESCE(%s(scenes.date - performers.birthdate), %s) FROM %s as performers INNER JOIN %s AS aggregation ON performers.id = aggregation.%s WHERE aggregation.%s = %s.id) %s", + aggregation, + fallback, + performerTable, + performersScenesTable, + performerIDColumn, + sceneIDColumn, + sceneTable, + getSortDirection(direction), + ) + case "studio": + query.joinSort(studioTable, "", "scenes.studio_id = studios.id") + add, agg := getSort("name", direction, studioTable) + query.sort += add + query.addGroupBy(agg...) + default: + add, agg := getSort(sort, direction, "scenes") + query.sort += add + query.addGroupBy(agg...) + } + + // Whatever the sorting, always use title/id as a final sort + query.sort += ", COALESCE(scenes.title, CAST(scenes.id as text)) COLLATE NATURAL_CI ASC" + query.addGroupBy("scenes.title", "scenes.id") + + return nil +} + +func (qb *SceneStore) SaveActivity(ctx context.Context, id int, resumeTime *float64, playDuration *float64) (bool, error) { + if err := qb.tableMgr.checkIDExists(ctx, id); err != nil { + return false, err + } + + record := goqu.Record{} + + if resumeTime != nil { + record["resume_time"] = resumeTime + } + + if playDuration != nil { + record["play_duration"] = goqu.L("play_duration + ?", playDuration) + } + + if len(record) > 0 { + if err := qb.tableMgr.updateByID(ctx, id, record); err != nil { + return false, err + } + } + + return true, nil +} + +func (qb *SceneStore) ResetActivity(ctx context.Context, id int, resetResume bool, resetDuration bool) (bool, error) { + if err := qb.tableMgr.checkIDExists(ctx, id); err != nil { + return false, err + } + + record := goqu.Record{} + + if resetResume { + record["resume_time"] = 0.0 + } + + if resetDuration { + record["play_duration"] = 0.0 + } + + if len(record) > 0 { + if err := qb.tableMgr.updateByID(ctx, id, record); err != nil { + return false, err + } + } + + return true, nil +} + +func (qb *SceneStore) GetURLs(ctx context.Context, sceneID int) ([]string, error) { + return scenesURLsTableMgr.get(ctx, sceneID) +} + +func (qb *SceneStore) GetCover(ctx context.Context, sceneID int) ([]byte, error) { + return qb.GetImage(ctx, sceneID, sceneCoverBlobColumn) +} + +func (qb *SceneStore) HasCover(ctx context.Context, sceneID int) (bool, error) { + return qb.HasImage(ctx, sceneID, sceneCoverBlobColumn) +} + +func (qb *SceneStore) UpdateCover(ctx context.Context, sceneID int, image []byte) error { + return qb.UpdateImage(ctx, sceneID, sceneCoverBlobColumn, image) +} + +func (qb *SceneStore) destroyCover(ctx context.Context, sceneID int) error { + return qb.DestroyImage(ctx, sceneID, sceneCoverBlobColumn) +} + +func (qb *SceneStore) AssignFiles(ctx context.Context, sceneID int, fileIDs []models.FileID) error { + // assuming a file can only be assigned to a single scene + if err := scenesFilesTableMgr.destroyJoins(ctx, fileIDs); err != nil { + return err + } + + // assign primary only if destination has no files + existingFileIDs, err := sceneRepository.files.get(ctx, sceneID) + if err != nil { + return err + } + + firstPrimary := len(existingFileIDs) == 0 + return scenesFilesTableMgr.insertJoins(ctx, sceneID, firstPrimary, fileIDs) +} + +func (qb *SceneStore) GetGroups(ctx context.Context, id int) (ret []models.GroupsScenes, err error) { + ret = []models.GroupsScenes{} + + if err := sceneRepository.groups.getAll(ctx, id, func(rows *sqlx.Rows) error { + var ms groupsScenesRow + if err := rows.StructScan(&ms); err != nil { + return err + } + + ret = append(ret, ms.resolve(id)) + return nil + }); err != nil { + return nil, err + } + + return ret, nil +} + +func (qb *SceneStore) AddFileID(ctx context.Context, id int, fileID models.FileID) error { + const firstPrimary = false + return scenesFilesTableMgr.insertJoins(ctx, id, firstPrimary, []models.FileID{fileID}) +} + +func (qb *SceneStore) GetPerformerIDs(ctx context.Context, id int) ([]int, error) { + return sceneRepository.performers.getIDs(ctx, id) +} + +func (qb *SceneStore) GetTagIDs(ctx context.Context, id int) ([]int, error) { + return sceneRepository.tags.getIDs(ctx, id) +} + +func (qb *SceneStore) GetGalleryIDs(ctx context.Context, id int) ([]int, error) { + return sceneRepository.galleries.getIDs(ctx, id) +} + +func (qb *SceneStore) AddGalleryIDs(ctx context.Context, sceneID int, galleryIDs []int) error { + return scenesGalleriesTableMgr.addJoins(ctx, sceneID, galleryIDs) +} + +func (qb *SceneStore) GetStashIDs(ctx context.Context, sceneID int) ([]models.StashID, error) { + return sceneRepository.stashIDs.get(ctx, sceneID) +} + +func (qb *SceneStore) FindDuplicates(ctx context.Context, distance int, durationDiff float64) ([][]*models.Scene, error) { + var dupeIds [][]int + if distance == 0 { + var ids []string + if err := dbWrapper.Select(ctx, &ids, findExactDuplicateQuery, durationDiff); err != nil { + return nil, err + } + + for _, id := range ids { + strIds := strings.Split(id, ",") + var sceneIds []int + for _, strId := range strIds { + if intId, err := strconv.Atoi(strId); err == nil { + sceneIds = sliceutil.AppendUnique(sceneIds, intId) + } + } + // filter out + if len(sceneIds) > 1 { + dupeIds = append(dupeIds, sceneIds) + } + } + } else { + var hashes []*utils.Phash + + if err := sceneRepository.queryFunc(ctx, findAllPhashesQuery, nil, false, func(rows *sqlx.Rows) error { + phash := utils.Phash{ + Bucket: -1, + Duration: -1, + } + if err := rows.StructScan(&phash); err != nil { + return err + } + + hashes = append(hashes, &phash) + return nil + }); err != nil { + return nil, err + } + + dupeIds = utils.FindDuplicates(hashes, distance, durationDiff) + } + + var duplicates [][]*models.Scene + for _, sceneIds := range dupeIds { + if scenes, err := qb.FindMany(ctx, sceneIds); err == nil { + duplicates = append(duplicates, scenes) + } + } + + sortByPath(duplicates) + + return duplicates, nil +} + +func sortByPath(scenes [][]*models.Scene) { + lessFunc := func(i int, j int) bool { + firstPathI := getFirstPath(scenes[i]) + firstPathJ := getFirstPath(scenes[j]) + return firstPathI < firstPathJ + } + sort.SliceStable(scenes, lessFunc) +} + +func getFirstPath(scenes []*models.Scene) string { + var firstPath string + for i, scene := range scenes { + if i == 0 || scene.Path < firstPath { + firstPath = scene.Path + } + } + return firstPath +} diff --git a/pkg/postgres/scene_filter.go b/pkg/postgres/scene_filter.go new file mode 100644 index 0000000000..1f0afb988e --- /dev/null +++ b/pkg/postgres/scene_filter.go @@ -0,0 +1,586 @@ +package postgres + +import ( + "context" + "fmt" + + "github.com/stashapp/stash/pkg/models" + "github.com/stashapp/stash/pkg/utils" +) + +type sceneFilterHandler struct { + sceneFilter *models.SceneFilterType +} + +func (qb *sceneFilterHandler) validate() error { + sceneFilter := qb.sceneFilter + if sceneFilter == nil { + return nil + } + + if err := validateFilterCombination(sceneFilter.OperatorFilter); err != nil { + return err + } + + if subFilter := sceneFilter.SubFilter(); subFilter != nil { + sqb := &sceneFilterHandler{sceneFilter: subFilter} + if err := sqb.validate(); err != nil { + return err + } + } + + return nil +} + +func (qb *sceneFilterHandler) handle(ctx context.Context, f *filterBuilder) { + sceneFilter := qb.sceneFilter + if sceneFilter == nil { + return + } + + if err := qb.validate(); err != nil { + f.setError(err) + return + } + + sf := sceneFilter.SubFilter() + if sf != nil { + sub := &sceneFilterHandler{sf} + handleSubFilter(ctx, sub, f, sceneFilter.OperatorFilter) + } + + f.handleCriterion(ctx, qb.criterionHandler()) +} + +func (qb *sceneFilterHandler) criterionHandler() criterionHandler { + sceneFilter := qb.sceneFilter + return compoundHandler{ + intCriterionHandler(sceneFilter.ID, "scenes.id", nil), + pathCriterionHandler(sceneFilter.Path, "folders.path", "files.basename", qb.addFoldersTable), + qb.fileCountCriterionHandler(sceneFilter.FileCount), + stringCriterionHandler(sceneFilter.Title, "scenes.title"), + stringCriterionHandler(sceneFilter.Code, "scenes.code"), + stringCriterionHandler(sceneFilter.Details, "scenes.details"), + stringCriterionHandler(sceneFilter.Director, "scenes.director"), + criterionHandlerFunc(func(ctx context.Context, f *filterBuilder) { + if sceneFilter.Oshash != nil { + qb.addSceneFilesTable(f) + f.addLeftJoin(fingerprintTable, "fingerprints_oshash", "scenes_files.file_id = fingerprints_oshash.file_id AND fingerprints_oshash.type = 'oshash'") + } + + stringCriterionHandler(sceneFilter.Oshash, "fingerprints_oshash.fingerprint")(ctx, f) + }), + + criterionHandlerFunc(func(ctx context.Context, f *filterBuilder) { + if sceneFilter.Checksum != nil { + qb.addSceneFilesTable(f) + f.addLeftJoin(fingerprintTable, "fingerprints_md5", "scenes_files.file_id = fingerprints_md5.file_id AND fingerprints_md5.type = 'md5'") + } + + stringCriterionHandler(sceneFilter.Checksum, "fingerprints_md5.fingerprint")(ctx, f) + }), + + criterionHandlerFunc(func(ctx context.Context, f *filterBuilder) { + if sceneFilter.Phash != nil { + // backwards compatibility + qb.phashDistanceCriterionHandler(&models.PhashDistanceCriterionInput{ + Value: sceneFilter.Phash.Value, + Modifier: sceneFilter.Phash.Modifier, + })(ctx, f) + } + }), + + qb.phashDistanceCriterionHandler(sceneFilter.PhashDistance), + + intCriterionHandler(sceneFilter.Rating100, "scenes.rating", nil), + qb.oCountCriterionHandler(sceneFilter.OCounter), + boolCriterionHandler(sceneFilter.Organized, "scenes.organized", nil), + + floatIntCriterionHandler(sceneFilter.Duration, "video_files.duration", qb.addVideoFilesTable), + resolutionCriterionHandler(sceneFilter.Resolution, "video_files.height", "video_files.width", qb.addVideoFilesTable), + orientationCriterionHandler(sceneFilter.Orientation, "video_files.height", "video_files.width", qb.addVideoFilesTable), + floatIntCriterionHandler(sceneFilter.Framerate, "ROUND(video_files.frame_rate)", qb.addVideoFilesTable), + intCriterionHandler(sceneFilter.Bitrate, "video_files.bit_rate", qb.addVideoFilesTable), + qb.codecCriterionHandler(sceneFilter.VideoCodec, "video_files.video_codec", qb.addVideoFilesTable), + qb.codecCriterionHandler(sceneFilter.AudioCodec, "video_files.audio_codec", qb.addVideoFilesTable), + + qb.hasMarkersCriterionHandler(sceneFilter.HasMarkers), + qb.isMissingCriterionHandler(sceneFilter.IsMissing), + qb.urlsCriterionHandler(sceneFilter.URL), + + criterionHandlerFunc(func(ctx context.Context, f *filterBuilder) { + if sceneFilter.StashID != nil { + sceneRepository.stashIDs.join(f, "scene_stash_ids", "scenes.id") + uuidCriterionHandler(sceneFilter.StashID, "scene_stash_ids.stash_id")(ctx, f) + } + }), + &stashIDCriterionHandler{ + c: sceneFilter.StashIDEndpoint, + stashIDRepository: &sceneRepository.stashIDs, + stashIDTableAs: "scene_stash_ids", + parentIDCol: "scenes.id", + }, + &stashIDsCriterionHandler{ + c: sceneFilter.StashIDsEndpoint, + stashIDRepository: &sceneRepository.stashIDs, + stashIDTableAs: "scene_stash_ids", + parentIDCol: "scenes.id", + }, + + boolCriterionHandler(sceneFilter.Interactive, "video_files.interactive", qb.addVideoFilesTable), + intCriterionHandler(sceneFilter.InteractiveSpeed, "video_files.interactive_speed", qb.addVideoFilesTable), + + qb.captionCriterionHandler(sceneFilter.Captions), + + floatIntCriterionHandler(sceneFilter.ResumeTime, "scenes.resume_time", nil), + floatIntCriterionHandler(sceneFilter.PlayDuration, "scenes.play_duration", nil), + qb.playCountCriterionHandler(sceneFilter.PlayCount), + criterionHandlerFunc(func(ctx context.Context, f *filterBuilder) { + if sceneFilter.LastPlayedAt != nil { + f.addLeftJoin( + fmt.Sprintf("(SELECT %s, MAX(%s) as last_played_at FROM %s GROUP BY %s)", sceneIDColumn, sceneViewDateColumn, scenesViewDatesTable, sceneIDColumn), + "scene_last_view", + fmt.Sprintf("scene_last_view.%s = scenes.id", sceneIDColumn), + ) + h := timestampCriterionHandler{sceneFilter.LastPlayedAt, "IFNULL(last_played_at, datetime(0))", nil} + h.handle(ctx, f) + } + }), + + qb.tagsCriterionHandler(sceneFilter.Tags), + qb.tagCountCriterionHandler(sceneFilter.TagCount), + qb.performersCriterionHandler(sceneFilter.Performers), + qb.performerCountCriterionHandler(sceneFilter.PerformerCount), + studioCriterionHandler(sceneTable, sceneFilter.Studios), + + qb.groupsCriterionHandler(sceneFilter.Groups), + qb.moviesCriterionHandler(sceneFilter.Movies), + + qb.galleriesCriterionHandler(sceneFilter.Galleries), + qb.performerTagsCriterionHandler(sceneFilter.PerformerTags), + qb.performerFavoriteCriterionHandler(sceneFilter.PerformerFavorite), + qb.performerAgeCriterionHandler(sceneFilter.PerformerAge), + qb.phashDuplicatedCriterionHandler(sceneFilter.Duplicated, qb.addSceneFilesTable), + &dateCriterionHandler{sceneFilter.Date, "scenes.date", nil}, + ×tampCriterionHandler{sceneFilter.CreatedAt, "scenes.created_at", nil}, + ×tampCriterionHandler{sceneFilter.UpdatedAt, "scenes.updated_at", nil}, + + &relatedFilterHandler{ + relatedIDCol: "scenes_galleries.gallery_id", + relatedRepo: galleryRepository.repository, + relatedHandler: &galleryFilterHandler{sceneFilter.GalleriesFilter}, + joinFn: func(f *filterBuilder) { + sceneRepository.galleries.innerJoin(f, "", "scenes.id") + }, + }, + + &relatedFilterHandler{ + relatedIDCol: "performers_join.performer_id", + relatedRepo: performerRepository.repository, + relatedHandler: &performerFilterHandler{sceneFilter.PerformersFilter}, + joinFn: func(f *filterBuilder) { + sceneRepository.performers.innerJoin(f, "performers_join", "scenes.id") + }, + }, + + &relatedFilterHandler{ + relatedIDCol: "scenes.studio_id", + relatedRepo: studioRepository.repository, + relatedHandler: &studioFilterHandler{sceneFilter.StudiosFilter}, + }, + + &relatedFilterHandler{ + relatedIDCol: "scene_tag.tag_id", + relatedRepo: tagRepository.repository, + relatedHandler: &tagFilterHandler{sceneFilter.TagsFilter}, + joinFn: func(f *filterBuilder) { + sceneRepository.tags.innerJoin(f, "scene_tag", "scenes.id") + }, + }, + + &relatedFilterHandler{ + relatedIDCol: "groups_scenes.group_id", + relatedRepo: groupRepository.repository, + relatedHandler: &groupFilterHandler{sceneFilter.MoviesFilter}, + joinFn: func(f *filterBuilder) { + sceneRepository.groups.innerJoin(f, "", "scenes.id") + }, + }, + + &relatedFilterHandler{ + relatedIDCol: "files.id", + relatedRepo: fileRepository.repository, + relatedHandler: &fileFilterHandler{ + fileFilter: sceneFilter.FilesFilter, + isRelated: true, + }, + joinFn: func(f *filterBuilder) { + qb.addFilesTable(f) + qb.addFoldersTable(f) + }, + // don't use a subquery; join directly + directJoin: true, + }, + + &relatedFilterHandler{ + relatedIDCol: "scene_markers.id", + relatedRepo: sceneMarkerRepository.repository, + relatedHandler: &sceneMarkerFilterHandler{sceneFilter.MarkersFilter}, + joinFn: func(f *filterBuilder) { + f.addInnerJoin("scene_markers", "", "scenes.id") + }, + }, + } +} + +func (qb *sceneFilterHandler) addSceneFilesTable(f *filterBuilder) { + f.addLeftJoin(scenesFilesTable, "", "scenes_files.scene_id = scenes.id AND scenes_files.\"primary\" = true") +} + +func (qb *sceneFilterHandler) addFilesTable(f *filterBuilder) { + qb.addSceneFilesTable(f) + f.addLeftJoin(fileTable, "", "scenes_files.file_id = files.id") +} + +func (qb *sceneFilterHandler) addFoldersTable(f *filterBuilder) { + qb.addFilesTable(f) + f.addLeftJoin(folderTable, "", "files.parent_folder_id = folders.id") +} + +func (qb *sceneFilterHandler) addVideoFilesTable(f *filterBuilder) { + qb.addSceneFilesTable(f) + f.addLeftJoin(videoFileTable, "", "video_files.file_id = scenes_files.file_id") +} + +func (qb *sceneFilterHandler) playCountCriterionHandler(count *models.IntCriterionInput) criterionHandlerFunc { + h := countCriterionHandlerBuilder{ + primaryTable: sceneTable, + joinTable: scenesViewDatesTable, + primaryFK: sceneIDColumn, + } + + return h.handler(count) +} + +func (qb *sceneFilterHandler) oCountCriterionHandler(count *models.IntCriterionInput) criterionHandlerFunc { + h := countCriterionHandlerBuilder{ + primaryTable: sceneTable, + joinTable: scenesODatesTable, + primaryFK: sceneIDColumn, + } + + return h.handler(count) +} + +func (qb *sceneFilterHandler) fileCountCriterionHandler(fileCount *models.IntCriterionInput) criterionHandlerFunc { + h := countCriterionHandlerBuilder{ + primaryTable: sceneTable, + joinTable: scenesFilesTable, + primaryFK: sceneIDColumn, + } + + return h.handler(fileCount) +} + +func (qb *sceneFilterHandler) phashDuplicatedCriterionHandler(duplicatedFilter *models.PHashDuplicationCriterionInput, addJoinFn func(f *filterBuilder)) criterionHandlerFunc { + return func(ctx context.Context, f *filterBuilder) { + // TODO: Wishlist item: Implement Distance matching + if duplicatedFilter != nil { + if addJoinFn != nil { + addJoinFn(f) + } + + var v string + if *duplicatedFilter.Duplicated { + v = ">" + } else { + v = "=" + } + + f.addInnerJoin("(SELECT file_id FROM files_fingerprints INNER JOIN (SELECT fingerprint FROM files_fingerprints WHERE type = 'phash' GROUP BY fingerprint HAVING COUNT (fingerprint) "+v+" 1) dupes on files_fingerprints.fingerprint = dupes.fingerprint)", "scph", "scenes_files.file_id = scph.file_id") + } + } +} + +func (qb *sceneFilterHandler) codecCriterionHandler(codec *models.StringCriterionInput, codecColumn string, addJoinFn func(f *filterBuilder)) criterionHandlerFunc { + return func(ctx context.Context, f *filterBuilder) { + if codec != nil { + if addJoinFn != nil { + addJoinFn(f) + } + + stringCriterionHandler(codec, codecColumn)(ctx, f) + } + } +} + +func (qb *sceneFilterHandler) hasMarkersCriterionHandler(hasMarkers *string) criterionHandlerFunc { + return func(ctx context.Context, f *filterBuilder) { + if hasMarkers != nil { + f.addLeftJoin("scene_markers", "", "scene_markers.scene_id = scenes.id") + if *hasMarkers == "true" { + f.addHaving("count(scene_markers.scene_id) > 0") + } else { + f.addWhere("scene_markers.id IS NULL") + } + } + } +} + +func (qb *sceneFilterHandler) isMissingCriterionHandler(isMissing *string) criterionHandlerFunc { + return func(ctx context.Context, f *filterBuilder) { + if isMissing != nil && *isMissing != "" { + switch *isMissing { + case "url": + scenesURLsTableMgr.join(f, "", "scenes.id") + f.addWhere("scene_urls.url IS NULL") + case "galleries": + sceneRepository.galleries.join(f, "galleries_join", "scenes.id") + f.addWhere("galleries_join.scene_id IS NULL") + case "studio": + f.addWhere("scenes.studio_id IS NULL") + case "movie", "group": + sceneRepository.groups.join(f, "groups_join", "scenes.id") + f.addWhere("groups_join.scene_id IS NULL") + case "performers": + sceneRepository.performers.join(f, "performers_join", "scenes.id") + f.addWhere("performers_join.scene_id IS NULL") + case "date": + f.addWhere(`scenes.date IS NULL`) + case "tags": + sceneRepository.tags.join(f, "tags_join", "scenes.id") + f.addWhere("tags_join.scene_id IS NULL") + case "stash_id": + sceneRepository.stashIDs.join(f, "scene_stash_ids", "scenes.id") + f.addWhere("scene_stash_ids.scene_id IS NULL") + case "phash": + qb.addSceneFilesTable(f) + f.addLeftJoin(fingerprintTable, "fingerprints_phash", "scenes_files.file_id = fingerprints_phash.file_id AND fingerprints_phash.type = 'phash'") + f.addWhere("fingerprints_phash.fingerprint IS NULL") + case "cover": + f.addWhere("scenes.cover_blob IS NULL") + default: + f.addWhere("(scenes." + *isMissing + " IS NULL OR TRIM(CAST(scenes." + *isMissing + " AS TEXT)) = '')") + } + } + } +} + +func (qb *sceneFilterHandler) urlsCriterionHandler(url *models.StringCriterionInput) criterionHandlerFunc { + h := stringListCriterionHandlerBuilder{ + primaryTable: sceneTable, + primaryFK: sceneIDColumn, + joinTable: scenesURLsTable, + stringColumn: sceneURLColumn, + addJoinTable: func(f *filterBuilder) { + scenesURLsTableMgr.join(f, "", "scenes.id") + }, + } + + return h.handler(url) +} + +func (qb *sceneFilterHandler) getMultiCriterionHandlerBuilder(foreignTable, joinTable, foreignFK string, addJoinsFunc func(f *filterBuilder)) multiCriterionHandlerBuilder { + return multiCriterionHandlerBuilder{ + primaryTable: sceneTable, + foreignTable: foreignTable, + joinTable: joinTable, + primaryFK: sceneIDColumn, + foreignFK: foreignFK, + addJoinsFunc: addJoinsFunc, + } +} + +func (qb *sceneFilterHandler) captionCriterionHandler(captions *models.StringCriterionInput) criterionHandlerFunc { + h := stringListCriterionHandlerBuilder{ + primaryTable: sceneTable, + primaryFK: sceneIDColumn, + joinTable: videoCaptionsTable, + stringColumn: captionCodeColumn, + addJoinTable: func(f *filterBuilder) { + qb.addSceneFilesTable(f) + f.addLeftJoin(videoCaptionsTable, "", "video_captions.file_id = scenes_files.file_id") + }, + excludeHandler: func(f *filterBuilder, criterion *models.StringCriterionInput) { + excludeClause := `scenes.id NOT IN ( + SELECT scenes_files.scene_id from scenes_files + INNER JOIN video_captions on video_captions.file_id = scenes_files.file_id + WHERE video_captions.language_code ILIKE ? + )` + f.addWhere(excludeClause, criterion.Value) + + // TODO - should we also exclude null values? + }, + } + + return h.handler(captions) +} + +func (qb *sceneFilterHandler) tagsCriterionHandler(tags *models.HierarchicalMultiCriterionInput) criterionHandlerFunc { + h := joinedHierarchicalMultiCriterionHandlerBuilder{ + primaryTable: sceneTable, + foreignTable: tagTable, + foreignFK: "tag_id", + + relationsTable: "tags_relations", + joinAs: "scene_tag", + joinTable: scenesTagsTable, + primaryFK: sceneIDColumn, + } + + return h.handler(tags) +} + +func (qb *sceneFilterHandler) tagCountCriterionHandler(tagCount *models.IntCriterionInput) criterionHandlerFunc { + h := countCriterionHandlerBuilder{ + primaryTable: sceneTable, + joinTable: scenesTagsTable, + primaryFK: sceneIDColumn, + } + + return h.handler(tagCount) +} + +func (qb *sceneFilterHandler) performersCriterionHandler(performers *models.MultiCriterionInput) criterionHandlerFunc { + h := joinedMultiCriterionHandlerBuilder{ + primaryTable: sceneTable, + joinTable: performersScenesTable, + joinAs: "performers_join", + primaryFK: sceneIDColumn, + foreignFK: performerIDColumn, + + addJoinTable: func(f *filterBuilder) { + sceneRepository.performers.join(f, "performers_join", "scenes.id") + }, + } + + return h.handler(performers) +} + +func (qb *sceneFilterHandler) performerCountCriterionHandler(performerCount *models.IntCriterionInput) criterionHandlerFunc { + h := countCriterionHandlerBuilder{ + primaryTable: sceneTable, + joinTable: performersScenesTable, + primaryFK: sceneIDColumn, + } + + return h.handler(performerCount) +} + +func (qb *sceneFilterHandler) performerFavoriteCriterionHandler(performerfavorite *bool) criterionHandlerFunc { + return func(ctx context.Context, f *filterBuilder) { + if performerfavorite != nil { + f.addLeftJoin("performers_scenes", "", "scenes.id = performers_scenes.scene_id") + + if *performerfavorite { + // contains at least one favorite + f.addLeftJoin("performers", "", "performers.id = performers_scenes.performer_id") + f.addWhere("performers.favorite = true") + } else { + // contains zero favorites + f.addLeftJoin(`(SELECT performers_scenes.scene_id as id FROM performers_scenes +JOIN performers ON performers.id = performers_scenes.performer_id +GROUP BY performers_scenes.scene_id HAVING SUM(performers.favorite) = false)`, "nofaves", "scenes.id = nofaves.id") + f.addWhere("performers_scenes.scene_id IS NULL OR nofaves.id IS NOT NULL") + } + } + } +} + +func (qb *sceneFilterHandler) performerAgeCriterionHandler(performerAge *models.IntCriterionInput) criterionHandlerFunc { + return func(ctx context.Context, f *filterBuilder) { + if performerAge != nil { + f.addInnerJoin("performers_scenes", "", "scenes.id = performers_scenes.scene_id") + f.addInnerJoin("performers", "", "performers_scenes.performer_id = performers.id") + + f.addWhere("scenes.date != '' AND performers.birthdate != ''") + f.addWhere("scenes.date IS NOT NULL AND performers.birthdate IS NOT NULL") + + ageCalc := "EXTRACT(YEAR FROM AGE(scenes.date, performers.birthdate))" + whereClause, args := getIntWhereClause(ageCalc, performerAge.Modifier, performerAge.Value, performerAge.Value2) + f.addWhere(whereClause, args...) + } + } +} + +// legacy handler +func (qb *sceneFilterHandler) moviesCriterionHandler(movies *models.MultiCriterionInput) criterionHandlerFunc { + addJoinsFunc := func(f *filterBuilder) { + sceneRepository.groups.join(f, "", "scenes.id") + f.addLeftJoin("groups", "", "groups_scenes.group_id = groups.id") + } + h := qb.getMultiCriterionHandlerBuilder(groupTable, groupsScenesTable, "group_id", addJoinsFunc) + return h.handler(movies) +} + +func (qb *sceneFilterHandler) groupsCriterionHandler(groups *models.HierarchicalMultiCriterionInput) criterionHandlerFunc { + h := joinedHierarchicalMultiCriterionHandlerBuilder{ + primaryTable: sceneTable, + foreignTable: groupTable, + foreignFK: "group_id", + + relationsTable: groupRelationsTable, + parentFK: "containing_id", + childFK: "sub_id", + joinAs: "scene_group", + joinTable: groupsScenesTable, + primaryFK: sceneIDColumn, + } + + return h.handler(groups) +} + +func (qb *sceneFilterHandler) galleriesCriterionHandler(galleries *models.MultiCriterionInput) criterionHandlerFunc { + addJoinsFunc := func(f *filterBuilder) { + sceneRepository.galleries.join(f, "", "scenes.id") + f.addLeftJoin("galleries", "", "scenes_galleries.gallery_id = galleries.id") + } + h := qb.getMultiCriterionHandlerBuilder(galleryTable, scenesGalleriesTable, "gallery_id", addJoinsFunc) + return h.handler(galleries) +} + +func (qb *sceneFilterHandler) performerTagsCriterionHandler(tags *models.HierarchicalMultiCriterionInput) criterionHandler { + return &joinedPerformerTagsHandler{ + criterion: tags, + primaryTable: sceneTable, + joinTable: performersScenesTable, + joinPrimaryKey: sceneIDColumn, + } +} + +func (qb *sceneFilterHandler) phashDistanceCriterionHandler(phashDistance *models.PhashDistanceCriterionInput) criterionHandlerFunc { + return func(ctx context.Context, f *filterBuilder) { + if phashDistance != nil { + qb.addSceneFilesTable(f) + f.addLeftJoin(fingerprintTable, "fingerprints_phash", "scenes_files.file_id = fingerprints_phash.file_id AND fingerprints_phash.type = 'phash'") + + value, _ := utils.StringToPhash(phashDistance.Value) + distance := 0 + if phashDistance.Distance != nil { + distance = *phashDistance.Distance + } + + if distance == 0 { + // use the default handler + intCriterionHandler(&models.IntCriterionInput{ + Value: int(value), + Modifier: phashDistance.Modifier, + }, "CAST(fingerprints_phash.fingerprint AS bigint)", nil)(ctx, f) + } + + switch { + case phashDistance.Modifier == models.CriterionModifierEquals && distance > 0: + // needed to avoid a type mismatch + f.addWhere("phash_distance(CAST(fingerprints_phash.fingerprint AS bigint), ?) < ?", value, distance) + case phashDistance.Modifier == models.CriterionModifierNotEquals && distance > 0: + // needed to avoid a type mismatch + f.addWhere("phash_distance(CAST(fingerprints_phash.fingerprint AS bigint), ?) > ?", value, distance) + default: + intCriterionHandler(&models.IntCriterionInput{ + Value: int(value), + Modifier: phashDistance.Modifier, + }, "CAST(fingerprints_phash.fingerprint AS bigint)", nil)(ctx, f) + } + } + } +} diff --git a/pkg/postgres/scene_marker.go b/pkg/postgres/scene_marker.go new file mode 100644 index 0000000000..33d0ec878e --- /dev/null +++ b/pkg/postgres/scene_marker.go @@ -0,0 +1,481 @@ +package postgres + +import ( + "context" + "database/sql" + "errors" + "fmt" + "slices" + + "github.com/doug-martin/goqu/v9" + "github.com/doug-martin/goqu/v9/exp" + "github.com/jmoiron/sqlx" + "gopkg.in/guregu/null.v4" + + "github.com/stashapp/stash/pkg/models" +) + +const ( + sceneMarkerTable = "scene_markers" + sceneMarkersTagsTable = "scene_markers_tags" + sceneMarkerIDColumn = "scene_marker_id" +) + +const countSceneMarkersForTagQuery = ` +SELECT scene_markers.id FROM scene_markers +LEFT JOIN scene_markers_tags as tags_join on tags_join.scene_marker_id = scene_markers.id +WHERE tags_join.tag_id = ? OR scene_markers.primary_tag_id = ? +GROUP BY scene_markers.id +` + +type sceneMarkerRow struct { + ID int `db:"id" goqu:"skipinsert"` + Title string `db:"title"` // TODO: make db schema (and gql schema) nullable + Seconds float64 `db:"seconds"` + PrimaryTagID int `db:"primary_tag_id"` + SceneID int `db:"scene_id"` + CreatedAt Timestamp `db:"created_at"` + UpdatedAt Timestamp `db:"updated_at"` + EndSeconds null.Float `db:"end_seconds"` +} + +func (r *sceneMarkerRow) fromSceneMarker(o models.SceneMarker) { + r.ID = o.ID + r.Title = o.Title + r.Seconds = o.Seconds + if o.EndSeconds != nil { + r.EndSeconds = null.FloatFrom(*o.EndSeconds) + } + r.PrimaryTagID = o.PrimaryTagID + r.SceneID = o.SceneID + r.CreatedAt = Timestamp{Timestamp: o.CreatedAt} + r.UpdatedAt = Timestamp{Timestamp: o.UpdatedAt} +} + +func (r *sceneMarkerRow) resolve() *models.SceneMarker { + ret := &models.SceneMarker{ + ID: r.ID, + Title: r.Title, + Seconds: r.Seconds, + EndSeconds: r.EndSeconds.Ptr(), + PrimaryTagID: r.PrimaryTagID, + SceneID: r.SceneID, + CreatedAt: r.CreatedAt.Timestamp, + UpdatedAt: r.UpdatedAt.Timestamp, + } + + return ret +} + +type sceneMarkerRowRecord struct { + updateRecord +} + +func (r *sceneMarkerRowRecord) fromPartial(o models.SceneMarkerPartial) { + // TODO: replace with setNullString after schema is made nullable + // r.setNullString("title", o.Title) + // saves a null input as the empty string + if o.Title.Set { + r.set("title", o.Title.Value) + } + r.setFloat64("seconds", o.Seconds) + r.setNullFloat64("end_seconds", o.EndSeconds) + r.setInt("primary_tag_id", o.PrimaryTagID) + r.setInt("scene_id", o.SceneID) + r.setTimestamp("created_at", o.CreatedAt) + r.setTimestamp("updated_at", o.UpdatedAt) +} + +type sceneMarkerRepositoryType struct { + repository + + scenes repository + tags joinRepository +} + +var ( + sceneMarkerRepository = sceneMarkerRepositoryType{ + repository: repository{ + tableName: sceneMarkerTable, + idColumn: idColumn, + }, + scenes: repository{ + tableName: sceneTable, + idColumn: idColumn, + }, + tags: joinRepository{ + repository: repository{ + tableName: sceneMarkersTagsTable, + idColumn: sceneMarkerIDColumn, + }, + fkColumn: tagIDColumn, + }, + } +) + +type SceneMarkerStore struct{} + +func NewSceneMarkerStore() *SceneMarkerStore { + return &SceneMarkerStore{} +} + +func (qb *SceneMarkerStore) table() exp.IdentifierExpression { + return sceneMarkerTableMgr.table +} + +func (qb *SceneMarkerStore) selectDataset() *goqu.SelectDataset { + return dialect.From(qb.table()).Select(qb.table().All()) +} + +func (qb *SceneMarkerStore) Create(ctx context.Context, newObject *models.SceneMarker) error { + var r sceneMarkerRow + r.fromSceneMarker(*newObject) + + id, err := sceneMarkerTableMgr.insertID(ctx, r) + if err != nil { + return err + } + + updated, err := qb.find(ctx, id) + if err != nil { + return fmt.Errorf("finding after create: %w", err) + } + + *newObject = *updated + + return nil +} + +func (qb *SceneMarkerStore) UpdatePartial(ctx context.Context, id int, partial models.SceneMarkerPartial) (*models.SceneMarker, error) { + r := sceneMarkerRowRecord{ + updateRecord{ + Record: make(exp.Record), + }, + } + + r.fromPartial(partial) + + if len(r.Record) > 0 { + if err := sceneMarkerTableMgr.updateByID(ctx, id, r.Record); err != nil { + return nil, err + } + } + + if partial.TagIDs != nil { + if err := sceneMarkersTagsTableMgr.modifyJoins(ctx, id, partial.TagIDs.IDs, partial.TagIDs.Mode); err != nil { + return nil, fmt.Errorf("modifying scene marker tags: %w", err) + } + } + + return qb.find(ctx, id) +} + +func (qb *SceneMarkerStore) Update(ctx context.Context, updatedObject *models.SceneMarker) error { + var r sceneMarkerRow + r.fromSceneMarker(*updatedObject) + + if err := sceneMarkerTableMgr.updateByID(ctx, updatedObject.ID, r); err != nil { + return err + } + + return nil +} + +func (qb *SceneMarkerStore) Destroy(ctx context.Context, id int) error { + return sceneMarkerRepository.destroyExisting(ctx, []int{id}) +} + +// returns nil, nil if not found +func (qb *SceneMarkerStore) Find(ctx context.Context, id int) (*models.SceneMarker, error) { + ret, err := qb.find(ctx, id) + if errors.Is(err, sql.ErrNoRows) { + return nil, nil + } + return ret, err +} + +func (qb *SceneMarkerStore) FindMany(ctx context.Context, ids []int) ([]*models.SceneMarker, error) { + ret := make([]*models.SceneMarker, len(ids)) + + if len(ids) == 0 { + return ret, nil + } + + table := qb.table() + q := qb.selectDataset().Prepared(true).Where(table.Col(idColumn).In(ids)) + unsorted, err := qb.getMany(ctx, q) + if err != nil { + return nil, err + } + + for _, s := range unsorted { + i := slices.Index(ids, s.ID) + ret[i] = s + } + + for i := range ret { + if ret[i] == nil { + return nil, fmt.Errorf("scene marker with id %d not found", ids[i]) + } + } + + return ret, nil +} + +// returns nil, sql.ErrNoRows if not found +func (qb *SceneMarkerStore) find(ctx context.Context, id int) (*models.SceneMarker, error) { + q := qb.selectDataset().Where(sceneMarkerTableMgr.byID(id)) + + ret, err := qb.get(ctx, q) + if err != nil { + return nil, err + } + + return ret, nil +} + +// returns nil, sql.ErrNoRows if not found +func (qb *SceneMarkerStore) get(ctx context.Context, q *goqu.SelectDataset) (*models.SceneMarker, error) { + ret, err := qb.getMany(ctx, q) + if err != nil { + return nil, err + } + + if len(ret) == 0 { + return nil, sql.ErrNoRows + } + + return ret[0], nil +} + +func (qb *SceneMarkerStore) getMany(ctx context.Context, q *goqu.SelectDataset) ([]*models.SceneMarker, error) { + const single = false + var ret []*models.SceneMarker + if err := queryFunc(ctx, q, single, func(r *sqlx.Rows) error { + var f sceneMarkerRow + if err := r.StructScan(&f); err != nil { + return err + } + + s := f.resolve() + + ret = append(ret, s) + return nil + }); err != nil { + return nil, err + } + + return ret, nil +} + +func (qb *SceneMarkerStore) FindBySceneID(ctx context.Context, sceneID int) ([]*models.SceneMarker, error) { + query := ` + SELECT scene_markers.* FROM scene_markers + WHERE scene_markers.scene_id = ? + GROUP BY scene_markers.id + ORDER BY scene_markers.seconds ASC + ` + args := []interface{}{sceneID} + return qb.querySceneMarkers(ctx, query, args) +} + +func (qb *SceneMarkerStore) CountByTagID(ctx context.Context, tagID int) (int, error) { + args := []interface{}{tagID, tagID} + return sceneMarkerRepository.runCountQuery(ctx, sceneMarkerRepository.buildCountQuery(countSceneMarkersForTagQuery), args) +} + +func (qb *SceneMarkerStore) GetMarkerStrings(ctx context.Context, q *string, sort *string) ([]*models.MarkerStringsResultType, error) { + query := "SELECT count(*) as `count`, scene_markers.id as id, scene_markers.title as title FROM scene_markers" + if q != nil { + query += " WHERE title ILIKE '%" + *q + "%'" + } + query += " GROUP BY title" + if sort != nil && *sort == "count" { + query += " ORDER BY `count` DESC" + } else { + query += " ORDER BY title ASC" + } + var args []interface{} + return qb.queryMarkerStringsResultType(ctx, query, args) +} + +func (qb *SceneMarkerStore) Wall(ctx context.Context, q *string) ([]*models.SceneMarker, error) { + s := "" + if q != nil { + s = *q + } + + table := qb.table() + qq := qb.selectDataset().Prepared(true).Where(table.Col("title").ILike("%" + s + "%")).Order(goqu.L("RANDOM()").Asc()).Limit(80) + return qb.getMany(ctx, qq) +} + +func (qb *SceneMarkerStore) makeQuery(ctx context.Context, sceneMarkerFilter *models.SceneMarkerFilterType, findFilter *models.FindFilterType) (*queryBuilder, error) { + if sceneMarkerFilter == nil { + sceneMarkerFilter = &models.SceneMarkerFilterType{} + } + if findFilter == nil { + findFilter = &models.FindFilterType{} + } + + query := sceneMarkerRepository.newQuery() + distinctIDs(&query, sceneMarkerTable) + + if q := findFilter.Q; q != nil && *q != "" { + query.join(sceneTable, "", "scenes.id = scene_markers.scene_id") + query.join(tagTable, "", "scene_markers.primary_tag_id = tags.id") + searchColumns := []string{"scene_markers.title", "scenes.title", "tags.name"} + query.parseQueryString(searchColumns, *q) + } + + filter := filterBuilderFromHandler(ctx, &sceneMarkerFilterHandler{ + sceneMarkerFilter: sceneMarkerFilter, + }) + + if err := query.addFilter(filter); err != nil { + return nil, err + } + + if err := qb.setSceneMarkerSort(&query, findFilter); err != nil { + return nil, err + } + query.pagination = getPagination(findFilter) + + return &query, nil +} + +func (qb *SceneMarkerStore) Query(ctx context.Context, sceneMarkerFilter *models.SceneMarkerFilterType, findFilter *models.FindFilterType) ([]*models.SceneMarker, int, error) { + query, err := qb.makeQuery(ctx, sceneMarkerFilter, findFilter) + if err != nil { + return nil, 0, err + } + + idsResult, countResult, err := query.executeFind(ctx) + if err != nil { + return nil, 0, err + } + + sceneMarkers, err := qb.FindMany(ctx, idsResult) + if err != nil { + return nil, 0, err + } + + return sceneMarkers, countResult, nil +} + +func (qb *SceneMarkerStore) QueryCount(ctx context.Context, sceneMarkerFilter *models.SceneMarkerFilterType, findFilter *models.FindFilterType) (int, error) { + query, err := qb.makeQuery(ctx, sceneMarkerFilter, findFilter) + if err != nil { + return 0, err + } + + return query.executeCount(ctx) +} + +var sceneMarkerSortOptions = sortOptions{ + "created_at", + "id", + "title", + "random", + "scene_id", + "scenes_updated_at", + "seconds", + "updated_at", + "duration", +} + +func (qb *SceneMarkerStore) setSceneMarkerSort(query *queryBuilder, findFilter *models.FindFilterType) error { + sort := findFilter.GetSort("title") + direction := findFilter.GetDirection() + + // CVE-2024-32231 - ensure sort is in the list of allowed sorts + if err := sceneMarkerSortOptions.validateSort(sort); err != nil { + return err + } + + switch sort { + case "scenes_updated_at": + sort = "updated_at" + query.joinSort(sceneTable, "", "scenes.id = scene_markers.scene_id") + add, agg := getSort(sort, direction, sceneTable) + query.sort += add + query.addGroupBy(agg...) + case "title": + query.joinSort(tagTable, "", "scene_markers.primary_tag_id = tags.id") + query.sort += " ORDER BY COALESCE(NULLIF(scene_markers.title,''), tags.name) COLLATE NATURAL_CI " + direction + query.addGroupBy("scene_markers.title", "tags.name") + case "duration": + sort = "(scene_markers.end_seconds - scene_markers.seconds)" + add, _ := getSort(sort, direction, sceneMarkerTable) + query.sort += add + query.addGroupBy("scene_markers.end_seconds", "scene_markers.seconds") + default: + add, agg := getSort(sort, direction, sceneMarkerTable) + query.sort += add + query.addGroupBy(agg...) + } + + query.sort += ", scene_markers.scene_id ASC, scene_markers.seconds ASC" + return nil +} + +func (qb *SceneMarkerStore) querySceneMarkers(ctx context.Context, query string, args []interface{}) ([]*models.SceneMarker, error) { + const single = false + var ret []*models.SceneMarker + if err := sceneMarkerRepository.queryFunc(ctx, query, args, single, func(r *sqlx.Rows) error { + var f sceneMarkerRow + if err := r.StructScan(&f); err != nil { + return err + } + + s := f.resolve() + + ret = append(ret, s) + return nil + }); err != nil { + return nil, err + } + + return ret, nil +} + +func (qb *SceneMarkerStore) queryMarkerStringsResultType(ctx context.Context, query string, args []interface{}) ([]*models.MarkerStringsResultType, error) { + rows, err := dbWrapper.Queryx(ctx, query, args...) + if err != nil && !errors.Is(err, sql.ErrNoRows) { + return nil, err + } + defer rows.Close() + + markerStrings := make([]*models.MarkerStringsResultType, 0) + for rows.Next() { + markerString := models.MarkerStringsResultType{} + if err := rows.StructScan(&markerString); err != nil { + return nil, err + } + markerStrings = append(markerStrings, &markerString) + } + + if err := rows.Err(); err != nil { + return nil, err + } + + return markerStrings, nil +} + +func (qb *SceneMarkerStore) GetTagIDs(ctx context.Context, id int) ([]int, error) { + return sceneMarkerRepository.tags.getIDs(ctx, id) +} + +func (qb *SceneMarkerStore) UpdateTags(ctx context.Context, id int, tagIDs []int) error { + // Delete the existing joins and then create new ones + return sceneMarkerRepository.tags.replace(ctx, id, tagIDs) +} + +func (qb *SceneMarkerStore) Count(ctx context.Context) (int, error) { + q := dialect.Select(goqu.COUNT("*")).From(qb.table()) + return count(ctx, q) +} + +func (qb *SceneMarkerStore) All(ctx context.Context) ([]*models.SceneMarker, error) { + return qb.getMany(ctx, qb.selectDataset()) +} diff --git a/pkg/postgres/scene_marker_filter.go b/pkg/postgres/scene_marker_filter.go new file mode 100644 index 0000000000..f88b5f53e6 --- /dev/null +++ b/pkg/postgres/scene_marker_filter.go @@ -0,0 +1,206 @@ +package postgres + +import ( + "context" + "fmt" + + "github.com/stashapp/stash/pkg/models" +) + +type sceneMarkerFilterHandler struct { + sceneMarkerFilter *models.SceneMarkerFilterType +} + +func (qb *sceneMarkerFilterHandler) validate() error { + return nil +} + +func (qb *sceneMarkerFilterHandler) handle(ctx context.Context, f *filterBuilder) { + sceneMarkerFilter := qb.sceneMarkerFilter + if sceneMarkerFilter == nil { + return + } + + if err := qb.validate(); err != nil { + f.setError(err) + return + } + + f.handleCriterion(ctx, qb.criterionHandler()) +} + +func (qb *sceneMarkerFilterHandler) joinScenes(f *filterBuilder) { + sceneMarkerRepository.scenes.innerJoin(f, "", "scene_markers.scene_id") +} + +func (qb *sceneMarkerFilterHandler) criterionHandler() criterionHandler { + sceneMarkerFilter := qb.sceneMarkerFilter + return compoundHandler{ + qb.tagIDCriterionHandler(sceneMarkerFilter.TagID), + qb.tagsCriterionHandler(sceneMarkerFilter.Tags), + qb.sceneTagsCriterionHandler(sceneMarkerFilter.SceneTags), + qb.performersCriterionHandler(sceneMarkerFilter.Performers), + qb.scenesCriterionHandler(sceneMarkerFilter.Scenes), + floatCriterionHandler(sceneMarkerFilter.Duration, "COALESCE(scene_markers.end_seconds - scene_markers.seconds, NULL)", nil), + ×tampCriterionHandler{sceneMarkerFilter.CreatedAt, "scene_markers.created_at", nil}, + ×tampCriterionHandler{sceneMarkerFilter.UpdatedAt, "scene_markers.updated_at", nil}, + &dateCriterionHandler{sceneMarkerFilter.SceneDate, "scenes.date", qb.joinScenes}, + ×tampCriterionHandler{sceneMarkerFilter.SceneCreatedAt, "scenes.created_at", qb.joinScenes}, + ×tampCriterionHandler{sceneMarkerFilter.SceneUpdatedAt, "scenes.updated_at", qb.joinScenes}, + + &relatedFilterHandler{ + relatedIDCol: "scenes.id", + relatedRepo: sceneRepository.repository, + relatedHandler: &sceneFilterHandler{sceneMarkerFilter.SceneFilter}, + joinFn: func(f *filterBuilder) { + qb.joinScenes(f) + }, + }, + } +} + +func (qb *sceneMarkerFilterHandler) tagIDCriterionHandler(tagID *string) criterionHandlerFunc { + return func(ctx context.Context, f *filterBuilder) { + if tagID != nil { + f.addLeftJoin("scene_markers_tags", "", "scene_markers_tags.scene_marker_id = scene_markers.id") + + f.addWhere("(scene_markers.primary_tag_id = ? OR scene_markers_tags.tag_id = ?)", *tagID, *tagID) + } + } +} + +func (qb *sceneMarkerFilterHandler) tagsCriterionHandler(criterion *models.HierarchicalMultiCriterionInput) criterionHandlerFunc { + return func(ctx context.Context, f *filterBuilder) { + if criterion != nil { + tags := criterion.CombineExcludes() + + if tags.Modifier == models.CriterionModifierIsNull || tags.Modifier == models.CriterionModifierNotNull { + var notClause string + if tags.Modifier == models.CriterionModifierNotNull { + notClause = "NOT" + } + + f.addLeftJoin("scene_markers_tags", "", "scene_markers.id = scene_markers_tags.scene_marker_id") + + f.addWhere(fmt.Sprintf("%s scene_markers_tags.tag_id IS NULL", notClause)) + return + } + + if tags.Modifier == models.CriterionModifierEquals && tags.Depth != nil && *tags.Depth != 0 { + f.setError(fmt.Errorf("depth is not supported for equals modifier for marker tag filtering")) + return + } + + if len(tags.Value) == 0 && len(tags.Excludes) == 0 { + return + } + + if len(tags.Value) > 0 { + valuesClause, err := getHierarchicalValues(ctx, tags.Value, tagTable, "tags_relations", "parent_id", "child_id", tags.Depth, false) + if err != nil { + f.setError(err) + return + } + + f.addWith(`marker_tags AS ( + SELECT mt.scene_marker_id, t.column1 AS root_tag_id FROM scene_markers_tags mt + INNER JOIN (` + valuesClause + `) t ON t.column2 = mt.tag_id + UNION + SELECT m.id, t.column1 FROM scene_markers m + INNER JOIN (` + valuesClause + `) t ON t.column2 = m.primary_tag_id + )`) + + f.addLeftJoin("marker_tags", "", "marker_tags.scene_marker_id = scene_markers.id") + + switch tags.Modifier { + case models.CriterionModifierEquals: + // includes only the provided ids + f.addWhere("marker_tags.root_tag_id IS NOT NULL") + tagsLen := len(tags.Value) + f.addHaving(fmt.Sprintf("count(distinct marker_tags.root_tag_id) = %d", tagsLen)) + // decrement by one to account for primary tag id + f.addWhere("(SELECT COUNT(*) FROM scene_markers_tags s WHERE s.scene_marker_id = scene_markers.id) = ?", tagsLen-1) + case models.CriterionModifierNotEquals: + f.setError(fmt.Errorf("not equals modifier is not supported for scene marker tags")) + default: + addHierarchicalConditionClauses(f, tags, "marker_tags", "root_tag_id") + } + } + + if len(criterion.Excludes) > 0 { + valuesClause, err := getHierarchicalValues(ctx, tags.Excludes, tagTable, "tags_relations", "parent_id", "child_id", tags.Depth, true) + if err != nil { + f.setError(err) + return + } + + clause := "scene_markers.id NOT IN (SELECT scene_markers_tags.scene_marker_id FROM scene_markers_tags WHERE scene_markers_tags.tag_id IN (SELECT column2 FROM %s))" + f.addWhere(fmt.Sprintf(clause, valuesClause)) + + f.addWhere(fmt.Sprintf("scene_markers.primary_tag_id NOT IN (SELECT column2 FROM %s)", valuesClause)) + } + } + } +} + +func (qb *sceneMarkerFilterHandler) sceneTagsCriterionHandler(tags *models.HierarchicalMultiCriterionInput) criterionHandlerFunc { + return func(ctx context.Context, f *filterBuilder) { + if tags != nil { + f.addLeftJoin("scenes_tags", "", "scene_markers.scene_id = scenes_tags.scene_id") + + h := joinedHierarchicalMultiCriterionHandlerBuilder{ + primaryTable: "scene_markers", + primaryKey: sceneIDColumn, + foreignTable: tagTable, + foreignFK: tagIDColumn, + + relationsTable: "tags_relations", + joinTable: "scenes_tags", + joinAs: "marker_scenes_tags", + primaryFK: sceneIDColumn, + } + + h.handler(tags).handle(ctx, f) + } + } +} + +func (qb *sceneMarkerFilterHandler) performersCriterionHandler(performers *models.MultiCriterionInput) criterionHandlerFunc { + h := joinedMultiCriterionHandlerBuilder{ + primaryTable: sceneTable, + joinTable: performersScenesTable, + joinAs: "performers_join", + primaryFK: sceneIDColumn, + foreignFK: performerIDColumn, + + addJoinTable: func(f *filterBuilder) { + f.addLeftJoin(performersScenesTable, "performers_join", "performers_join.scene_id = scene_markers.scene_id") + }, + } + + handler := h.handler(performers) + return func(ctx context.Context, f *filterBuilder) { + if performers == nil { + return + } + + // Make sure scenes is included, otherwise excludes filter fails + qb.joinScenes(f) + handler(ctx, f) + } +} + +func (qb *sceneMarkerFilterHandler) scenesCriterionHandler(scenes *models.MultiCriterionInput) criterionHandlerFunc { + addJoinsFunc := func(f *filterBuilder) { + f.addLeftJoin(sceneTable, "markers_scenes", "markers_scenes.id = scene_markers.scene_id") + } + h := multiCriterionHandlerBuilder{ + primaryTable: sceneMarkerTable, + foreignTable: "markers_scenes", + joinTable: "", + primaryFK: sceneIDColumn, + foreignFK: sceneIDColumn, + addJoinsFunc: addJoinsFunc, + } + return h.handler(scenes) +} diff --git a/pkg/postgres/sql.go b/pkg/postgres/sql.go new file mode 100644 index 0000000000..caee4170e5 --- /dev/null +++ b/pkg/postgres/sql.go @@ -0,0 +1,392 @@ +package postgres + +import ( + "fmt" + "math/rand" + "regexp" + "strconv" + "strings" + "time" + + "github.com/stashapp/stash/pkg/models" +) + +func selectAll(tableName string) string { + idColumn := getColumn(tableName, "*") + return "SELECT " + idColumn + " FROM " + tableName + " " +} + +func distinctIDs(qb *queryBuilder, tableName string) { + columnId := getColumn(tableName, "id") + qb.addColumn(columnId) + qb.addGroupBy(columnId) + qb.from = tableName +} + +func selectIDs(qb *queryBuilder, tableName string) { + columnId := getColumn(tableName, "id") + qb.addColumn(columnId) + qb.from = tableName +} + +func getColumn(tableName string, columnName string) string { + return tableName + "." + columnName +} + +func getPagination(findFilter *models.FindFilterType) *queryPagination { + if findFilter == nil { + panic("nil find filter for pagination") + } + + if findFilter.IsGetAll() { + return nil + } + + return &queryPagination{ + page: findFilter.GetPage(), + perPage: findFilter.GetPageSize(), + } +} + +func getPaginationSQL(pag *queryPagination) string { + // Find all + if pag == nil || pag.perPage < 0 { + return " " + } + + // Counting + if pag.perPage == 0 { + return " LIMIT 1 OFFSET 0 " + } + + var page = (pag.page - 1) * pag.perPage + return " LIMIT " + strconv.Itoa(pag.perPage) + " OFFSET " + strconv.Itoa(page) + " " +} + +const randomSeedPrefix = "random_" // prefix for random sort + +type sortOptions []string + +func (o sortOptions) validateSort(sort string) error { + if strings.HasPrefix(sort, randomSeedPrefix) { + // seed as a parameter from the UI + seedStr := sort[len(randomSeedPrefix):] + _, err := strconv.ParseUint(seedStr, 10, 64) + if err != nil { + return fmt.Errorf("invalid random seed: %s", seedStr) + } + return nil + } + + for _, v := range o { + if v == sort { + return nil + } + } + + return fmt.Errorf("invalid sort: %s", sort) +} + +func getSortDirection(direction string) string { + if direction != "ASC" && direction != "DESC" { + return "ASC NULLS LAST" + } else { + return direction + " NULLS LAST" + } +} +func getSort(sort string, direction string, tableName string) (string, []string) { + direction = getSortDirection(direction) + + switch { + case strings.HasSuffix(sort, "_count"): + var relationTableName = strings.TrimSuffix(sort, "_count") // TODO: pluralize? + colName := getColumn(relationTableName, "id") + return " ORDER BY COUNT(distinct " + colName + ") " + direction, nil + case strings.Compare(sort, "filesize") == 0: + colName := getColumn(tableName, "size") + return " ORDER BY " + colName + " " + direction, []string{colName} + case strings.HasPrefix(sort, randomSeedPrefix): + // seed as a parameter from the UI + seedStr := sort[len(randomSeedPrefix):] + seed, err := strconv.ParseUint(seedStr, 10, 64) + if err != nil { + // fallback to a random seed + seed = rand.Uint64() + } + return getRandomSort(tableName, direction, seed), nil + case strings.Compare(sort, "random") == 0: + return getRandomSort(tableName, direction, rand.Uint64()), nil + default: + colName := getColumn(tableName, sort) + if strings.Contains(sort, ".") { + colName = sort + } + if strings.Compare(sort, "name") == 0 { + return " ORDER BY " + colName + " COLLATE NATURAL_CI " + direction, []string{colName} + } + if strings.Compare(sort, "title") == 0 { + return " ORDER BY " + colName + " COLLATE NATURAL_CI " + direction, []string{colName} + } + + return " ORDER BY " + colName + " " + direction, []string{colName} + } +} + +func getRandomSort(tableName string, direction string, seed uint64) string { + // cap seed at 10^8 + seed %= 1e8 + + colName := "CAST(" + getColumn(tableName, "id") + " AS DECIMAL)" + + // https://stackoverflow.com/questions/21949795#comment33255354_21949859 + // p1 := 52959209 + // p2 := 1047483763 + // p3 := 2147483647 + // n := + // ORDER BY ((n+seed)*(n+seed)*p1 + (n+seed)*p2) % p3 + // since sqlite converts overflowing numbers to reals, a custom db function that uses uints with overflow should be faster, + // however in practice the overhead of calling a custom function vastly outweighs the benefits + return fmt.Sprintf(" ORDER BY mod((%[1]s + %[2]d) * (%[1]s + %[2]d) * 52959209 + (%[1]s + %[2]d) * 1047483763, 2147483647) %[3]s", colName, seed, direction) +} + +func getCountSort(primaryTable, joinTable, primaryFK, direction string) string { + return fmt.Sprintf(" ORDER BY (SELECT COUNT(*) FROM %s AS sort WHERE sort.%s = %s.id) %s", joinTable, primaryFK, primaryTable, getSortDirection(direction)) +} + +// getStringSearchClause returns a sqlClause for searching strings in the provided columns. +// It is used for includes and excludes string criteria. +func getStringSearchClause(columns []string, q string, not bool) sqlClause { + var likeClauses []string + var args []interface{} + + notStr := "" + binaryType := " OR " + if not { + notStr = " NOT" + binaryType = " AND " + } + q = strings.TrimSpace(q) + trimmedQuery := strings.Trim(q, "\"") + + if trimmedQuery == q { + q = regexp.MustCompile(`\s+`).ReplaceAllString(q, " ") + queryWords := strings.Split(q, " ") + // Search for any word + for _, word := range queryWords { + for _, column := range columns { + likeClauses = append(likeClauses, column+notStr+" ILIKE ?") + args = append(args, "%"+word+"%") + } + } + } else { + // Search the exact query + for _, column := range columns { + likeClauses = append(likeClauses, column+notStr+" ILIKE ?") + args = append(args, "%"+trimmedQuery+"%") + } + } + likes := strings.Join(likeClauses, binaryType) + + return makeClause("("+likes+")", args...) +} + +func getEnumSearchClause(column string, enumVals []string, not bool) sqlClause { + var args []interface{} + + notStr := "" + if not { + notStr = " NOT" + } + + clause := fmt.Sprintf("(%s%s IN %s)", column, notStr, getInBinding(len(enumVals))) + for _, enumVal := range enumVals { + args = append(args, enumVal) + } + + return makeClause(clause, args...) +} + +func getInBinding(length int) string { + bindings := strings.Repeat("?, ", length) + bindings = strings.TrimRight(bindings, ", ") + return "(" + bindings + ")" +} + +func getIntCriterionWhereClause(column string, input models.IntCriterionInput) (string, []interface{}) { + return getIntWhereClause(column, input.Modifier, input.Value, input.Value2) +} + +func getIntWhereClause(column string, modifier models.CriterionModifier, value int, upper *int) (string, []interface{}) { + if upper == nil { + u := 0 + upper = &u + } + + args := []interface{}{value, *upper} + return getNumericWhereClause(column, modifier, args) +} + +func getFloatCriterionWhereClause(column string, input models.FloatCriterionInput) (string, []interface{}) { + return getFloatWhereClause(column, input.Modifier, input.Value, input.Value2) +} + +func getFloatWhereClause(column string, modifier models.CriterionModifier, value float64, upper *float64) (string, []interface{}) { + if upper == nil { + u := 0.0 + upper = &u + } + + args := []interface{}{value, *upper} + return getNumericWhereClause(column, modifier, args) +} + +func getNumericWhereClause(column string, modifier models.CriterionModifier, args []interface{}) (string, []interface{}) { + singleArgs := args[0:1] + + switch modifier { + case models.CriterionModifierIsNull: + return fmt.Sprintf("%s IS NULL", column), nil + case models.CriterionModifierNotNull: + return fmt.Sprintf("%s IS NOT NULL", column), nil + case models.CriterionModifierEquals: + return fmt.Sprintf("%s = ?", column), singleArgs + case models.CriterionModifierNotEquals: + return fmt.Sprintf("%s != ?", column), singleArgs + case models.CriterionModifierBetween: + return fmt.Sprintf("%s BETWEEN ? AND ?", column), args + case models.CriterionModifierNotBetween: + return fmt.Sprintf("%s NOT BETWEEN ? AND ?", column), args + case models.CriterionModifierLessThan: + return fmt.Sprintf("%s < ?", column), singleArgs + case models.CriterionModifierGreaterThan: + return fmt.Sprintf("%s > ?", column), singleArgs + } + + panic("unsupported numeric modifier type " + modifier) +} + +func getDateCriterionWhereClause(column string, input models.DateCriterionInput) (string, []interface{}) { + return getDateWhereClause(column, input.Modifier, input.Value, input.Value2) +} + +func getDateWhereClause(column string, modifier models.CriterionModifier, value string, upper *string) (string, []interface{}) { + if upper == nil { + u := time.Now().AddDate(0, 0, 1).Format(time.RFC3339) + upper = &u + } + + args := []interface{}{value} + betweenArgs := []interface{}{value, *upper} + + switch modifier { + case models.CriterionModifierIsNull: + return fmt.Sprintf("(%s IS NULL OR %s = '')", column, column), nil + case models.CriterionModifierNotNull: + return fmt.Sprintf("(%s IS NOT NULL AND %s != '')", column, column), nil + case models.CriterionModifierEquals: + return fmt.Sprintf("%s = ?", column), args + case models.CriterionModifierNotEquals: + return fmt.Sprintf("%s != ?", column), args + case models.CriterionModifierBetween: + return fmt.Sprintf("%s BETWEEN ? AND ?", column), betweenArgs + case models.CriterionModifierNotBetween: + return fmt.Sprintf("%s NOT BETWEEN ? AND ?", column), betweenArgs + case models.CriterionModifierLessThan: + return fmt.Sprintf("%s < ?", column), args + case models.CriterionModifierGreaterThan: + return fmt.Sprintf("%s > ?", column), args + } + + panic("unsupported date modifier type") +} + +func getTimestampCriterionWhereClause(column string, input models.TimestampCriterionInput) (string, []interface{}) { + return getTimestampWhereClause(column, input.Modifier, input.Value, input.Value2) +} + +func getTimestampWhereClause(column string, modifier models.CriterionModifier, value string, upper *string) (string, []interface{}) { + if upper == nil { + u := time.Now().AddDate(0, 0, 1).Format(time.RFC3339) + upper = &u + } + + args := []interface{}{value} + betweenArgs := []interface{}{value, *upper} + + switch modifier { + case models.CriterionModifierIsNull: + return fmt.Sprintf("%s IS NULL", column), nil + case models.CriterionModifierNotNull: + return fmt.Sprintf("%s IS NOT NULL", column), nil + case models.CriterionModifierEquals: + return fmt.Sprintf("%s = ?", column), args + case models.CriterionModifierNotEquals: + return fmt.Sprintf("%s != ?", column), args + case models.CriterionModifierBetween: + return fmt.Sprintf("%s BETWEEN ? AND ?", column), betweenArgs + case models.CriterionModifierNotBetween: + return fmt.Sprintf("%s NOT BETWEEN ? AND ?", column), betweenArgs + case models.CriterionModifierLessThan: + return fmt.Sprintf("%s < ?", column), args + case models.CriterionModifierGreaterThan: + return fmt.Sprintf("%s > ?", column), args + } + + panic("unsupported date modifier type") +} + +// returns where clause and having clause +func getMultiCriterionClause(primaryTable, foreignTable, joinTable, primaryFK, foreignFK string, criterion *models.MultiCriterionInput) (string, string) { + whereClause := "" + havingClause := "" + switch criterion.Modifier { + case models.CriterionModifierIncludes: + // includes any of the provided ids + if joinTable != "" { + whereClause = joinTable + "." + foreignFK + " IN " + getInBinding(len(criterion.Value)) + } else { + whereClause = foreignTable + ".id IN " + getInBinding(len(criterion.Value)) + } + case models.CriterionModifierIncludesAll: + // includes all of the provided ids + if joinTable != "" { + whereClause = joinTable + "." + foreignFK + " IN " + getInBinding(len(criterion.Value)) + havingClause = "count(distinct " + joinTable + "." + foreignFK + ") = " + strconv.Itoa(len(criterion.Value)) + } else { + whereClause = foreignTable + ".id IN " + getInBinding(len(criterion.Value)) + havingClause = "count(distinct " + foreignTable + ".id) = " + strconv.Itoa(len(criterion.Value)) + } + case models.CriterionModifierExcludes: + // excludes all of the provided ids + if joinTable != "" { + whereClause = primaryTable + ".id not in (select " + joinTable + "." + primaryFK + " from " + joinTable + " where " + joinTable + "." + foreignFK + " in " + getInBinding(len(criterion.Value)) + ")" + } else { + whereClause = "not exists (select s.id from " + primaryTable + " as s where s.id = " + primaryTable + ".id and s." + foreignFK + " in " + getInBinding(len(criterion.Value)) + ")" + } + } + + return whereClause, havingClause +} + +func getCountCriterionClause(primaryTable, joinTable, primaryFK string, criterion models.IntCriterionInput) (string, []interface{}) { + lhs := fmt.Sprintf("(SELECT COUNT(*) FROM %s s WHERE s.%s = %s.id)", joinTable, primaryFK, primaryTable) + return getIntCriterionWhereClause(lhs, criterion) +} + +func coalesce(column string) string { + return fmt.Sprintf("COALESCE(%s, '')", column) +} + +func like(v string) string { + return "%" + v + "%" +} + +type sqlTable string + +func (t sqlTable) Name() string { + return string(t) +} + +func (t sqlTable) Col(n string) string { + return fmt.Sprintf("%s.%s", string(t), n) +} diff --git a/pkg/postgres/studio.go b/pkg/postgres/studio.go new file mode 100644 index 0000000000..502f3551f2 --- /dev/null +++ b/pkg/postgres/studio.go @@ -0,0 +1,694 @@ +package postgres + +import ( + "context" + "database/sql" + "errors" + "fmt" + "slices" + + "github.com/doug-martin/goqu/v9" + "github.com/doug-martin/goqu/v9/exp" + "github.com/jmoiron/sqlx" + "gopkg.in/guregu/null.v4" + "gopkg.in/guregu/null.v4/zero" + + "github.com/stashapp/stash/pkg/models" + "github.com/stashapp/stash/pkg/studio" +) + +const ( + studioTable = "studios" + studioIDColumn = "studio_id" + + studioURLsTable = "studio_urls" + studioURLColumn = "url" + + studioAliasesTable = "studio_aliases" + studioAliasColumn = "alias" + studioParentIDColumn = "parent_id" + studioNameColumn = "name" + studioImageBlobColumn = "image_blob" + studiosTagsTable = "studios_tags" +) + +type studioRow struct { + ID int `db:"id" goqu:"skipinsert"` + Name zero.String `db:"name"` + ParentID null.Int `db:"parent_id,omitempty"` + CreatedAt Timestamp `db:"created_at"` + UpdatedAt Timestamp `db:"updated_at"` + // expressed as 1-100 + Rating null.Int `db:"rating"` + Favorite bool `db:"favorite"` + Details zero.String `db:"details"` + IgnoreAutoTag bool `db:"ignore_auto_tag"` + + // not used in resolutions or updates + ImageBlob zero.String `db:"image_blob"` +} + +func (r *studioRow) fromStudio(o models.Studio) { + r.ID = o.ID + r.Name = zero.StringFrom(o.Name) + r.ParentID = intFromPtr(o.ParentID) + r.CreatedAt = Timestamp{Timestamp: o.CreatedAt} + r.UpdatedAt = Timestamp{Timestamp: o.UpdatedAt} + r.Rating = intFromPtr(o.Rating) + r.Favorite = o.Favorite + r.Details = zero.StringFrom(o.Details) + r.IgnoreAutoTag = o.IgnoreAutoTag +} + +func (r *studioRow) resolve() *models.Studio { + ret := &models.Studio{ + ID: r.ID, + Name: r.Name.String, + ParentID: nullIntPtr(r.ParentID), + CreatedAt: r.CreatedAt.Timestamp.UTC(), + UpdatedAt: r.UpdatedAt.Timestamp.UTC(), + Rating: nullIntPtr(r.Rating), + Favorite: r.Favorite, + Details: r.Details.String, + IgnoreAutoTag: r.IgnoreAutoTag, + } + + return ret +} + +type studioRowRecord struct { + updateRecord +} + +func (r *studioRowRecord) fromPartial(o models.StudioPartial) { + r.setNullString("name", o.Name) + r.setNullInt("parent_id", o.ParentID) + r.setTimestamp("created_at", o.CreatedAt) + r.setTimestamp("updated_at", o.UpdatedAt) + r.setNullInt("rating", o.Rating) + r.setBool("favorite", o.Favorite) + r.setNullString("details", o.Details) + r.setBool("ignore_auto_tag", o.IgnoreAutoTag) +} + +type studioRepositoryType struct { + repository + + stashIDs stashIDRepository + tags joinRepository + + scenes repository + images repository + galleries repository +} + +var ( + studioRepository = studioRepositoryType{ + repository: repository{ + tableName: studioTable, + idColumn: idColumn, + }, + stashIDs: stashIDRepository{ + repository{ + tableName: "studio_stash_ids", + idColumn: studioIDColumn, + }, + }, + scenes: repository{ + tableName: sceneTable, + idColumn: studioIDColumn, + }, + images: repository{ + tableName: imageTable, + idColumn: studioIDColumn, + }, + galleries: repository{ + tableName: galleryTable, + idColumn: studioIDColumn, + }, + tags: joinRepository{ + repository: repository{ + tableName: studiosTagsTable, + idColumn: studioIDColumn, + }, + fkColumn: tagIDColumn, + foreignTable: tagTable, + orderBy: tagTableSortSQL, + }, + } +) + +type StudioStore struct { + blobJoinQueryBuilder + tagRelationshipStore + + tableMgr *table +} + +func NewStudioStore(blobStore *BlobStore) *StudioStore { + return &StudioStore{ + blobJoinQueryBuilder: blobJoinQueryBuilder{ + blobStore: blobStore, + joinTable: studioTable, + }, + tagRelationshipStore: tagRelationshipStore{ + idRelationshipStore: idRelationshipStore{ + joinTable: studiosTagsTableMgr, + }, + }, + + tableMgr: studioTableMgr, + } +} + +func (qb *StudioStore) table() exp.IdentifierExpression { + return qb.tableMgr.table +} + +func (qb *StudioStore) selectDataset() *goqu.SelectDataset { + return dialect.From(qb.table()).Select(qb.table().All()) +} + +func (qb *StudioStore) Create(ctx context.Context, newObject *models.Studio) error { + var err error + + var r studioRow + r.fromStudio(*newObject) + + id, err := qb.tableMgr.insertID(ctx, r) + if err != nil { + return err + } + + if newObject.Aliases.Loaded() { + if err := studio.ValidateAliases(ctx, id, newObject.Aliases.List(), qb); err != nil { + return err + } + + if err := studiosAliasesTableMgr.insertJoins(ctx, id, newObject.Aliases.List()); err != nil { + return err + } + } + + if newObject.URLs.Loaded() { + const startPos = 0 + if err := studiosURLsTableMgr.insertJoins(ctx, id, startPos, newObject.URLs.List()); err != nil { + return err + } + } + + if err := qb.tagRelationshipStore.createRelationships(ctx, id, newObject.TagIDs); err != nil { + return err + } + + if newObject.StashIDs.Loaded() { + if err := studiosStashIDsTableMgr.insertJoins(ctx, id, newObject.StashIDs.List()); err != nil { + return err + } + } + + updated, err := qb.find(ctx, id) + if err != nil { + return fmt.Errorf("finding after create: %w", err) + } + + *newObject = *updated + return nil +} + +func (qb *StudioStore) UpdatePartial(ctx context.Context, input models.StudioPartial) (*models.Studio, error) { + r := studioRowRecord{ + updateRecord{ + Record: make(exp.Record), + }, + } + + r.fromPartial(input) + + if len(r.Record) > 0 { + if err := qb.tableMgr.updateByID(ctx, input.ID, r.Record); err != nil { + return nil, err + } + } + + if input.Aliases != nil { + if err := studiosAliasesTableMgr.modifyJoins(ctx, input.ID, input.Aliases.Values, input.Aliases.Mode); err != nil { + return nil, err + } + } + + if input.URLs != nil { + if err := studiosURLsTableMgr.modifyJoins(ctx, input.ID, input.URLs.Values, input.URLs.Mode); err != nil { + return nil, err + } + } + + if err := qb.tagRelationshipStore.modifyRelationships(ctx, input.ID, input.TagIDs); err != nil { + return nil, err + } + + if input.StashIDs != nil { + if err := studiosStashIDsTableMgr.modifyJoins(ctx, input.ID, input.StashIDs.StashIDs, input.StashIDs.Mode); err != nil { + return nil, err + } + } + + return qb.Find(ctx, input.ID) +} + +// This is only used by the Import/Export functionality +func (qb *StudioStore) Update(ctx context.Context, updatedObject *models.Studio) error { + var r studioRow + r.fromStudio(*updatedObject) + + if err := qb.tableMgr.updateByID(ctx, updatedObject.ID, r); err != nil { + return err + } + + if updatedObject.Aliases.Loaded() { + if err := studiosAliasesTableMgr.replaceJoins(ctx, updatedObject.ID, updatedObject.Aliases.List()); err != nil { + return err + } + } + + if updatedObject.URLs.Loaded() { + if err := studiosURLsTableMgr.replaceJoins(ctx, updatedObject.ID, updatedObject.URLs.List()); err != nil { + return err + } + } + + if err := qb.tagRelationshipStore.replaceRelationships(ctx, updatedObject.ID, updatedObject.TagIDs); err != nil { + return err + } + + if updatedObject.StashIDs.Loaded() { + if err := studiosStashIDsTableMgr.replaceJoins(ctx, updatedObject.ID, updatedObject.StashIDs.List()); err != nil { + return err + } + } + + return nil +} + +func (qb *StudioStore) Destroy(ctx context.Context, id int) error { + // must handle image checksums manually + if err := qb.destroyImage(ctx, id); err != nil { + return err + } + + return studioRepository.destroyExisting(ctx, []int{id}) +} + +// returns nil, nil if not found +func (qb *StudioStore) Find(ctx context.Context, id int) (*models.Studio, error) { + ret, err := qb.find(ctx, id) + if errors.Is(err, sql.ErrNoRows) { + return nil, nil + } + return ret, err +} + +func (qb *StudioStore) FindMany(ctx context.Context, ids []int) ([]*models.Studio, error) { + ret := make([]*models.Studio, len(ids)) + + if len(ids) == 0 { + return ret, nil + } + + table := qb.table() + if err := batchExec(ids, defaultBatchSize, func(batch []int) error { + q := qb.selectDataset().Prepared(true).Where(table.Col(idColumn).In(batch)) + unsorted, err := qb.getMany(ctx, q) + if err != nil { + return err + } + + for _, s := range unsorted { + i := slices.Index(ids, s.ID) + ret[i] = s + } + + return nil + }); err != nil { + return nil, err + } + + for i := range ret { + if ret[i] == nil { + return nil, fmt.Errorf("studio with id %d not found", ids[i]) + } + } + + return ret, nil +} + +// returns nil, sql.ErrNoRows if not found +func (qb *StudioStore) find(ctx context.Context, id int) (*models.Studio, error) { + q := qb.selectDataset().Where(qb.tableMgr.byID(id)) + + ret, err := qb.get(ctx, q) + if err != nil { + return nil, err + } + + return ret, nil +} + +// returns nil, sql.ErrNoRows if not found +func (qb *StudioStore) get(ctx context.Context, q *goqu.SelectDataset) (*models.Studio, error) { + ret, err := qb.getMany(ctx, q) + if err != nil { + return nil, err + } + + if len(ret) == 0 { + return nil, sql.ErrNoRows + } + + return ret[0], nil +} + +func (qb *StudioStore) getMany(ctx context.Context, q *goqu.SelectDataset) ([]*models.Studio, error) { + const single = false + var ret []*models.Studio + if err := queryFunc(ctx, q, single, func(r *sqlx.Rows) error { + var f studioRow + if err := r.StructScan(&f); err != nil { + return err + } + + s := f.resolve() + + ret = append(ret, s) + return nil + }); err != nil { + return nil, err + } + + return ret, nil +} + +func (qb *StudioStore) findBySubquery(ctx context.Context, sq *goqu.SelectDataset) ([]*models.Studio, error) { + table := qb.table() + + q := qb.selectDataset().Where( + table.Col(idColumn).Eq( + sq, + ), + ) + + return qb.getMany(ctx, q) +} + +func (qb *StudioStore) FindChildren(ctx context.Context, id int) ([]*models.Studio, error) { + // SELECT studios.* FROM studios WHERE studios.parent_id = ? + table := qb.table() + sq := qb.selectDataset().Where(table.Col(studioParentIDColumn).Eq(id)) + ret, err := qb.getMany(ctx, sq) + + if err != nil { + return nil, err + } + + return ret, nil +} + +func (qb *StudioStore) FindBySceneID(ctx context.Context, sceneID int) (*models.Studio, error) { + // SELECT studios.* FROM studios JOIN scenes ON studios.id = scenes.studio_id WHERE scenes.id = ? LIMIT 1 + table := qb.table() + scenes := sceneTableMgr.table + sq := qb.selectDataset().Join( + scenes, goqu.On(table.Col(idColumn), scenes.Col(studioIDColumn)), + ).Where( + scenes.Col(idColumn), + ).Limit(1) + ret, err := qb.get(ctx, sq) + + if err != nil && !errors.Is(err, sql.ErrNoRows) { + return nil, err + } + + return ret, nil +} + +func (qb *StudioStore) FindByName(ctx context.Context, name string, nocase bool) (*models.Studio, error) { + // query := "SELECT * FROM studios WHERE name = ?" + // if nocase { + // query += " COLLATE NOCASE" + // } + // query += " LIMIT 1" + where := "name = ?" + if nocase { + where += " COLLATE NOCASE" + } + sq := qb.selectDataset().Prepared(true).Where(goqu.L(where, name)).Limit(1) + ret, err := qb.get(ctx, sq) + + if err != nil && !errors.Is(err, sql.ErrNoRows) { + return nil, err + } + + return ret, nil +} + +func (qb *StudioStore) FindByStashID(ctx context.Context, stashID models.StashID) ([]*models.Studio, error) { + sq := dialect.From(studiosStashIDsJoinTable).Select(studiosStashIDsJoinTable.Col(studioIDColumn)).Where( + studiosStashIDsJoinTable.Col("stash_id").Eq(stashID.StashID), + studiosStashIDsJoinTable.Col("endpoint").Eq(stashID.Endpoint), + ) + ret, err := qb.findBySubquery(ctx, sq) + + if err != nil { + return nil, fmt.Errorf("getting studios for stash ID %s: %w", stashID.StashID, err) + } + + return ret, nil +} + +func (qb *StudioStore) FindByStashIDStatus(ctx context.Context, hasStashID bool, stashboxEndpoint string) ([]*models.Studio, error) { + table := qb.table() + sq := dialect.From(table).LeftJoin( + studiosStashIDsJoinTable, + goqu.On(table.Col(idColumn).Eq(studiosStashIDsJoinTable.Col(studioIDColumn))), + ).Select(table.Col(idColumn)) + + if hasStashID { + sq = sq.Where( + studiosStashIDsJoinTable.Col("stash_id").IsNotNull(), + studiosStashIDsJoinTable.Col("endpoint").Eq(stashboxEndpoint), + ) + } else { + sq = sq.Where( + studiosStashIDsJoinTable.Col("stash_id").IsNull(), + ) + } + + ret, err := qb.findBySubquery(ctx, sq) + + if err != nil { + return nil, fmt.Errorf("getting studios for stash-box endpoint %s: %w", stashboxEndpoint, err) + } + + return ret, nil +} + +func (qb *StudioStore) Count(ctx context.Context) (int, error) { + q := dialect.Select(goqu.COUNT("*")).From(qb.table()) + return count(ctx, q) +} + +func (qb *StudioStore) All(ctx context.Context) ([]*models.Studio, error) { + table := qb.table() + return qb.getMany(ctx, qb.selectDataset().Order(table.Col(studioNameColumn).Asc())) +} + +func (qb *StudioStore) QueryForAutoTag(ctx context.Context, words []string) ([]*models.Studio, error) { + // TODO - Query needs to be changed to support queries of this type, and + // this method should be removed + table := qb.table() + sq := dialect.From(table).Select(table.Col(idColumn)).LeftJoin( + studiosAliasesJoinTable, + goqu.On(studiosAliasesJoinTable.Col(studioIDColumn).Eq(table.Col(idColumn))), + ) + + var whereClauses []exp.Expression + + for _, w := range words { + whereClauses = append(whereClauses, table.Col(studioNameColumn).ILike(w+"%")) + whereClauses = append(whereClauses, studiosAliasesJoinTable.Col("alias").ILike(w+"%")) + } + + sq = sq.Where( + goqu.Or(whereClauses...), + table.Col("ignore_auto_tag").IsFalse(), + ) + + ret, err := qb.findBySubquery(ctx, sq) + + if err != nil { + return nil, fmt.Errorf("getting studios for autotag: %w", err) + } + + return ret, nil +} + +func (qb *StudioStore) makeQuery(ctx context.Context, studioFilter *models.StudioFilterType, findFilter *models.FindFilterType) (*queryBuilder, error) { + if studioFilter == nil { + studioFilter = &models.StudioFilterType{} + } + if findFilter == nil { + findFilter = &models.FindFilterType{} + } + + query := studioRepository.newQuery() + distinctIDs(&query, studioTable) + + if q := findFilter.Q; q != nil && *q != "" { + query.join(studioAliasesTable, "", "studio_aliases.studio_id = studios.id") + searchColumns := []string{"studios.name", "studio_aliases.alias"} + query.parseQueryString(searchColumns, *q) + } + + filter := filterBuilderFromHandler(ctx, &studioFilterHandler{ + studioFilter: studioFilter, + }) + + if err := query.addFilter(filter); err != nil { + return nil, err + } + + var err error + var group []string + query.sort, group, err = qb.getStudioSort(findFilter) + if err != nil { + return nil, err + } + query.pagination = getPagination(findFilter) + query.addGroupBy(group...) + + return &query, nil +} + +func (qb *StudioStore) Query(ctx context.Context, studioFilter *models.StudioFilterType, findFilter *models.FindFilterType) ([]*models.Studio, int, error) { + query, err := qb.makeQuery(ctx, studioFilter, findFilter) + if err != nil { + return nil, 0, err + } + + idsResult, countResult, err := query.executeFind(ctx) + if err != nil { + return nil, 0, err + } + + studios, err := qb.FindMany(ctx, idsResult) + if err != nil { + return nil, 0, err + } + + return studios, countResult, nil +} + +func (qb *StudioStore) QueryCount(ctx context.Context, studioFilter *models.StudioFilterType, findFilter *models.FindFilterType) (int, error) { + query, err := qb.makeQuery(ctx, studioFilter, findFilter) + if err != nil { + return 0, err + } + + return query.executeCount(ctx) +} + +func (qb *StudioStore) sortByScenesDuration(direction string) string { + return fmt.Sprintf(` ORDER BY ( + SELECT COALESCE(SUM(video_files.duration), 0) + FROM %s + LEFT JOIN %s ON %s.%s = %s.id + LEFT JOIN video_files ON video_files.file_id = %s.file_id + WHERE %s.%s = %s.id + ) %s`, sceneTable, scenesFilesTable, scenesFilesTable, sceneIDColumn, sceneTable, scenesFilesTable, sceneTable, studioIDColumn, studioTable, getSortDirection(direction)) +} + +var studioSortOptions = sortOptions{ + "child_count", + "created_at", + "galleries_count", + "id", + "images_count", + "name", + "scenes_count", + "scenes_duration", + "random", + "rating", + "tag_count", + "updated_at", +} + +func (qb *StudioStore) getStudioSort(findFilter *models.FindFilterType) (string, []string, error) { + var sort string + var direction string + if findFilter == nil { + sort = "name" + direction = "ASC" + } else { + sort = findFilter.GetSort("name") + direction = findFilter.GetDirection() + } + + // CVE-2024-32231 - ensure sort is in the list of allowed sorts + if err := studioSortOptions.validateSort(sort); err != nil { + return "", nil, err + } + + group := []string{} + sortQuery := "" + switch sort { + case "tag_count": + sortQuery += getCountSort(studioTable, studiosTagsTable, studioIDColumn, direction) + case "scenes_count": + sortQuery += getCountSort(studioTable, sceneTable, studioIDColumn, direction) + case "scenes_duration": + sortQuery += qb.sortByScenesDuration(direction) + case "images_count": + sortQuery += getCountSort(studioTable, imageTable, studioIDColumn, direction) + case "galleries_count": + sortQuery += getCountSort(studioTable, galleryTable, studioIDColumn, direction) + case "child_count": + sortQuery += getCountSort(studioTable, studioTable, studioParentIDColumn, direction) + default: + var add string + add, group = getSort(sort, direction, "studios") + sortQuery += add + } + + // Whatever the sorting, always use name/id as a final sort + sortQuery += ", COALESCE(studios.name, CAST(studios.id as text)) COLLATE NATURAL_CI ASC" + group = append(group, "studios.name", "studios.id") + return sortQuery, group, nil +} + +func (qb *StudioStore) GetImage(ctx context.Context, studioID int) ([]byte, error) { + return qb.blobJoinQueryBuilder.GetImage(ctx, studioID, studioImageBlobColumn) +} + +func (qb *StudioStore) HasImage(ctx context.Context, studioID int) (bool, error) { + return qb.blobJoinQueryBuilder.HasImage(ctx, studioID, studioImageBlobColumn) +} + +func (qb *StudioStore) UpdateImage(ctx context.Context, studioID int, image []byte) error { + return qb.blobJoinQueryBuilder.UpdateImage(ctx, studioID, studioImageBlobColumn, image) +} + +func (qb *StudioStore) destroyImage(ctx context.Context, studioID int) error { + return qb.blobJoinQueryBuilder.DestroyImage(ctx, studioID, studioImageBlobColumn) +} + +func (qb *StudioStore) GetStashIDs(ctx context.Context, studioID int) ([]models.StashID, error) { + return studiosStashIDsTableMgr.get(ctx, studioID) +} + +func (qb *StudioStore) GetAliases(ctx context.Context, studioID int) ([]string, error) { + return studiosAliasesTableMgr.get(ctx, studioID) +} + +func (qb *StudioStore) GetURLs(ctx context.Context, studioID int) ([]string, error) { + return studiosURLsTableMgr.get(ctx, studioID) +} diff --git a/pkg/postgres/studio_filter.go b/pkg/postgres/studio_filter.go new file mode 100644 index 0000000000..ef5c704d5a --- /dev/null +++ b/pkg/postgres/studio_filter.go @@ -0,0 +1,252 @@ +package postgres + +import ( + "context" + + "github.com/stashapp/stash/pkg/models" +) + +type studioFilterHandler struct { + studioFilter *models.StudioFilterType +} + +func (qb *studioFilterHandler) validate() error { + studioFilter := qb.studioFilter + if studioFilter == nil { + return nil + } + + if err := validateFilterCombination(studioFilter.OperatorFilter); err != nil { + return err + } + + if subFilter := studioFilter.SubFilter(); subFilter != nil { + sqb := &studioFilterHandler{studioFilter: subFilter} + if err := sqb.validate(); err != nil { + return err + } + } + + return nil +} + +func (qb *studioFilterHandler) handle(ctx context.Context, f *filterBuilder) { + studioFilter := qb.studioFilter + if studioFilter == nil { + return + } + + if err := qb.validate(); err != nil { + f.setError(err) + return + } + + sf := studioFilter.SubFilter() + if sf != nil { + sub := &studioFilterHandler{sf} + handleSubFilter(ctx, sub, f, studioFilter.OperatorFilter) + } + + f.handleCriterion(ctx, qb.criterionHandler()) +} + +func (qb *studioFilterHandler) criterionHandler() criterionHandler { + studioFilter := qb.studioFilter + return compoundHandler{ + stringCriterionHandler(studioFilter.Name, studioTable+".name"), + stringCriterionHandler(studioFilter.Details, studioTable+".details"), + qb.urlsCriterionHandler(studioFilter.URL), + intCriterionHandler(studioFilter.Rating100, studioTable+".rating", nil), + boolCriterionHandler(studioFilter.Favorite, studioTable+".favorite", nil), + boolCriterionHandler(studioFilter.IgnoreAutoTag, studioTable+".ignore_auto_tag", nil), + + criterionHandlerFunc(func(ctx context.Context, f *filterBuilder) { + if studioFilter.StashID != nil { + studioRepository.stashIDs.join(f, "studio_stash_ids", "studios.id") + uuidCriterionHandler(studioFilter.StashID, "studio_stash_ids.stash_id")(ctx, f) + } + }), + &stashIDCriterionHandler{ + c: studioFilter.StashIDEndpoint, + stashIDRepository: &studioRepository.stashIDs, + stashIDTableAs: "studio_stash_ids", + parentIDCol: "studios.id", + }, + &stashIDsCriterionHandler{ + c: studioFilter.StashIDsEndpoint, + stashIDRepository: &studioRepository.stashIDs, + stashIDTableAs: "studio_stash_ids", + parentIDCol: "studios.id", + }, + + qb.isMissingCriterionHandler(studioFilter.IsMissing), + qb.tagCountCriterionHandler(studioFilter.TagCount), + qb.sceneCountCriterionHandler(studioFilter.SceneCount), + qb.imageCountCriterionHandler(studioFilter.ImageCount), + qb.galleryCountCriterionHandler(studioFilter.GalleryCount), + qb.parentCriterionHandler(studioFilter.Parents), + qb.aliasCriterionHandler(studioFilter.Aliases), + qb.tagsCriterionHandler(studioFilter.Tags), + qb.childCountCriterionHandler(studioFilter.ChildCount), + ×tampCriterionHandler{studioFilter.CreatedAt, studioTable + ".created_at", nil}, + ×tampCriterionHandler{studioFilter.UpdatedAt, studioTable + ".updated_at", nil}, + + &relatedFilterHandler{ + relatedIDCol: "scenes.id", + relatedRepo: sceneRepository.repository, + relatedHandler: &sceneFilterHandler{studioFilter.ScenesFilter}, + joinFn: func(f *filterBuilder) { + studioRepository.scenes.innerJoin(f, "", "studios.id") + }, + }, + + &relatedFilterHandler{ + relatedIDCol: "images.id", + relatedRepo: imageRepository.repository, + relatedHandler: &imageFilterHandler{studioFilter.ImagesFilter}, + joinFn: func(f *filterBuilder) { + studioRepository.images.innerJoin(f, "", "studios.id") + }, + }, + + &relatedFilterHandler{ + relatedIDCol: "galleries.id", + relatedRepo: galleryRepository.repository, + relatedHandler: &galleryFilterHandler{studioFilter.GalleriesFilter}, + joinFn: func(f *filterBuilder) { + studioRepository.galleries.innerJoin(f, "", "studios.id") + }, + }, + } +} + +func (qb *studioFilterHandler) isMissingCriterionHandler(isMissing *string) criterionHandlerFunc { + return func(ctx context.Context, f *filterBuilder) { + if isMissing != nil && *isMissing != "" { + switch *isMissing { + case "url": + studiosURLsTableMgr.join(f, "", "studios.id") + f.addWhere("studio_urls.url IS NULL") + case "image": + f.addWhere("studios.image_blob IS NULL") + case "stash_id": + studioRepository.stashIDs.join(f, "studio_stash_ids", "studios.id") + f.addWhere("studio_stash_ids.studio_id IS NULL") + default: + f.addWhere("(studios." + *isMissing + " IS NULL OR TRIM(CAST(studios." + *isMissing + " AS TEXT)) = '')") + } + } + } +} + +func (qb *studioFilterHandler) sceneCountCriterionHandler(sceneCount *models.IntCriterionInput) criterionHandlerFunc { + return func(ctx context.Context, f *filterBuilder) { + if sceneCount != nil { + f.addLeftJoin("scenes", "", "scenes.studio_id = studios.id") + clause, args := getIntCriterionWhereClause("count(distinct scenes.id)", *sceneCount) + + f.addHaving(clause, args...) + } + } +} + +func (qb *studioFilterHandler) imageCountCriterionHandler(imageCount *models.IntCriterionInput) criterionHandlerFunc { + return func(ctx context.Context, f *filterBuilder) { + if imageCount != nil { + f.addLeftJoin("images", "", "images.studio_id = studios.id") + clause, args := getIntCriterionWhereClause("count(distinct images.id)", *imageCount) + + f.addHaving(clause, args...) + } + } +} + +func (qb *studioFilterHandler) galleryCountCriterionHandler(galleryCount *models.IntCriterionInput) criterionHandlerFunc { + return func(ctx context.Context, f *filterBuilder) { + if galleryCount != nil { + f.addLeftJoin("galleries", "", "galleries.studio_id = studios.id") + clause, args := getIntCriterionWhereClause("count(distinct galleries.id)", *galleryCount) + + f.addHaving(clause, args...) + } + } +} + +func (qb *studioFilterHandler) tagCountCriterionHandler(tagCount *models.IntCriterionInput) criterionHandlerFunc { + h := countCriterionHandlerBuilder{ + primaryTable: studioTable, + joinTable: studiosTagsTable, + primaryFK: studioIDColumn, + } + + return h.handler(tagCount) +} + +func (qb *studioFilterHandler) parentCriterionHandler(parents *models.MultiCriterionInput) criterionHandlerFunc { + addJoinsFunc := func(f *filterBuilder) { + f.addLeftJoin("studios", "parent_studio", "parent_studio.id = studios.parent_id") + } + h := multiCriterionHandlerBuilder{ + primaryTable: studioTable, + foreignTable: "parent_studio", + joinTable: "", + primaryFK: studioIDColumn, + foreignFK: "parent_id", + addJoinsFunc: addJoinsFunc, + } + return h.handler(parents) +} + +func (qb *studioFilterHandler) aliasCriterionHandler(alias *models.StringCriterionInput) criterionHandlerFunc { + h := stringListCriterionHandlerBuilder{ + primaryTable: studioTable, + primaryFK: studioIDColumn, + joinTable: studioAliasesTable, + stringColumn: studioAliasColumn, + addJoinTable: func(f *filterBuilder) { + studiosAliasesTableMgr.join(f, "", "studios.id") + }, + } + + return h.handler(alias) +} + +func (qb *studioFilterHandler) urlsCriterionHandler(url *models.StringCriterionInput) criterionHandlerFunc { + h := stringListCriterionHandlerBuilder{ + primaryTable: studioTable, + primaryFK: studioIDColumn, + joinTable: studioURLsTable, + stringColumn: studioURLColumn, + addJoinTable: func(f *filterBuilder) { + studiosURLsTableMgr.join(f, "", "studios.id") + }, + } + + return h.handler(url) +} + +func (qb *studioFilterHandler) childCountCriterionHandler(childCount *models.IntCriterionInput) criterionHandlerFunc { + return func(ctx context.Context, f *filterBuilder) { + if childCount != nil { + f.addLeftJoin("studios", "children_count", "children_count.parent_id = studios.id") + clause, args := getIntCriterionWhereClause("count(distinct children_count.id)", *childCount) + + f.addHaving(clause, args...) + } + } +} + +func (qb *studioFilterHandler) tagsCriterionHandler(tags *models.HierarchicalMultiCriterionInput) criterionHandlerFunc { + h := joinedHierarchicalMultiCriterionHandlerBuilder{ + primaryTable: studioTable, + foreignTable: tagTable, + foreignFK: "tag_id", + + relationsTable: "tags_relations", + joinTable: studiosTagsTable, + joinAs: "studio_tag", + primaryFK: studioIDColumn, + } + + return h.handler(tags) +} diff --git a/pkg/postgres/table.go b/pkg/postgres/table.go new file mode 100644 index 0000000000..941760c8e5 --- /dev/null +++ b/pkg/postgres/table.go @@ -0,0 +1,1236 @@ +package postgres + +import ( + "context" + "database/sql" + "errors" + "fmt" + "time" + + "github.com/doug-martin/goqu/v9" + "github.com/doug-martin/goqu/v9/exp" + "github.com/jmoiron/sqlx" + "gopkg.in/guregu/null.v4" + + "github.com/stashapp/stash/pkg/logger" + "github.com/stashapp/stash/pkg/models" + "github.com/stashapp/stash/pkg/sliceutil" +) + +type table struct { + table exp.IdentifierExpression + idColumn exp.IdentifierExpression +} + +type NotFoundError struct { + ID int + Table string +} + +func (e *NotFoundError) Error() string { + return fmt.Sprintf("id %d does not exist in %s", e.ID, e.Table) +} + +func (t *table) insert(ctx context.Context, o interface{}) (sql.Result, error) { + q := dialect.Insert(t.table).Prepared(true).Rows(o) + ret, err := exec(ctx, q) + if err != nil { + return nil, fmt.Errorf("inserting into %s: %w", t.table.GetTable(), err) + } + + return ret, nil +} + +func (t *table) insertID(ctx context.Context, o interface{}) (int, error) { + q := dialect.Insert(t.table).Prepared(true).Rows(o).Returning(goqu.I("id")) + val, err := execID(ctx, q) + if err != nil { + return -1, fmt.Errorf("inserting into %s: %w", t.table.GetTable(), err) + } + + return int(*val), nil +} + +func (t *table) updateByID(ctx context.Context, id interface{}, o interface{}) error { + q := dialect.Update(t.table).Prepared(true).Set(o).Where(t.byID(id)) + + if _, err := exec(ctx, q); err != nil { + return fmt.Errorf("updating %s: %w", t.table.GetTable(), err) + } + + return nil +} + +func (t *table) byID(id interface{}) exp.Expression { + return t.idColumn.Eq(id) +} + +func (t *table) byIDInts(ids ...int) exp.Expression { + ii := make([]interface{}, len(ids)) + for i, id := range ids { + ii[i] = id + } + return t.idColumn.In(ii...) +} + +func (t *table) idExists(ctx context.Context, id interface{}) (bool, error) { + q := dialect.Select(goqu.COUNT("*")).From(t.table).Where(t.byID(id)) + + var count int + if err := querySimple(ctx, q, &count); err != nil { + return false, err + } + + return count == 1, nil +} + +func (t *table) checkIDExists(ctx context.Context, id int) error { + exists, err := t.idExists(ctx, id) + if err != nil { + return err + } + + if !exists { + return &NotFoundError{ID: id, Table: t.table.GetTable()} + } + + return nil +} + +func (t *table) destroyExisting(ctx context.Context, ids []int) error { + for _, id := range ids { + exists, err := t.idExists(ctx, id) + if err != nil { + return err + } + + if !exists { + return &NotFoundError{ + ID: id, + Table: t.table.GetTable(), + } + } + } + + return t.destroy(ctx, ids) +} + +func (t *table) destroy(ctx context.Context, ids []int) error { + q := dialect.Delete(t.table).Where(t.idColumn.In(ids)) + + if _, err := exec(ctx, q); err != nil { + return fmt.Errorf("destroying %s: %w", t.table.GetTable(), err) + } + + return nil +} + +func (t *table) join(j joiner, as string, parentIDCol string) { + tableName := t.table.GetTable() + tt := tableName + if as != "" { + tt = as + } + j.addLeftJoin(tableName, as, fmt.Sprintf("%s.%s = %s", tt, t.idColumn.GetCol(), parentIDCol)) +} + +// func (t *table) get(ctx context.Context, q *goqu.SelectDataset, dest interface{}) error { +// tx, err := getTx(ctx) +// if err != nil { +// return err +// } + +// sql, args, err := q.ToSQL() +// if err != nil { +// return fmt.Errorf("generating sql: %w", err) +// } + +// return tx.GetContext(ctx, dest, sql, args...) +// } + +type joinTable struct { + table + fkColumn exp.IdentifierExpression + + // required for ordering + foreignTable *table + orderBy exp.OrderedExpression +} + +func (t *joinTable) invert() *joinTable { + return &joinTable{ + table: table{ + table: t.table.table, + idColumn: t.fkColumn, + }, + fkColumn: t.table.idColumn, + foreignTable: t.foreignTable, + orderBy: t.orderBy, + } +} + +func (t *joinTable) get(ctx context.Context, id int) ([]int, error) { + q := dialect.Select(t.fkColumn).From(t.table.table).Where(t.idColumn.Eq(id)) + + if t.orderBy != nil { + if t.foreignTable != nil { + q = q.InnerJoin(t.foreignTable.table, goqu.On(t.foreignTable.idColumn.Eq(t.fkColumn))) + } + q = q.Order(t.orderBy) + } + + const single = false + var ret []int + if err := queryFunc(ctx, q, single, func(rows *sqlx.Rows) error { + var fk int + if err := rows.Scan(&fk); err != nil { + return err + } + + ret = append(ret, fk) + + return nil + }); err != nil { + return nil, fmt.Errorf("getting foreign keys from %s: %w", t.table.table.GetTable(), err) + } + + return ret, nil +} + +func (t *joinTable) insertJoins(ctx context.Context, id int, foreignIDs []int) error { + // manually create SQL so that we can prepare once + // ignore duplicates + q := fmt.Sprintf("INSERT INTO %s (%s, %s) VALUES (?, ?) ON CONFLICT (%[2]s, %s) DO NOTHING", t.table.table.GetTable(), t.idColumn.GetCol(), t.fkColumn.GetCol()) + + stmt, err := dbWrapper.Prepare(ctx, q) + if err != nil { + return err + } + defer stmt.Close() + + // eliminate duplicates + foreignIDs = sliceutil.AppendUniques(nil, foreignIDs) + + for _, fk := range foreignIDs { + if _, err := dbWrapper.ExecStmt(ctx, stmt, id, fk); err != nil { + return err + } + } + + return nil +} + +func (t *joinTable) replaceJoins(ctx context.Context, id int, foreignIDs []int) error { + if err := t.destroy(ctx, []int{id}); err != nil { + return err + } + + return t.insertJoins(ctx, id, foreignIDs) +} + +func (t *joinTable) addJoins(ctx context.Context, id int, foreignIDs []int) error { + // get existing foreign keys + fks, err := t.get(ctx, id) + if err != nil { + return err + } + + // only add foreign keys that are not already present + foreignIDs = sliceutil.Exclude(foreignIDs, fks) + return t.insertJoins(ctx, id, foreignIDs) +} + +func (t *joinTable) destroyJoins(ctx context.Context, id int, foreignIDs []int) error { + q := dialect.Delete(t.table.table).Where( + t.idColumn.Eq(id), + t.fkColumn.In(foreignIDs), + ) + + if _, err := exec(ctx, q); err != nil { + return fmt.Errorf("destroying %s: %w", t.table.table.GetTable(), err) + } + + return nil +} + +func (t *joinTable) modifyJoins(ctx context.Context, id int, foreignIDs []int, mode models.RelationshipUpdateMode) error { + switch mode { + case models.RelationshipUpdateModeSet: + return t.replaceJoins(ctx, id, foreignIDs) + case models.RelationshipUpdateModeAdd: + return t.addJoins(ctx, id, foreignIDs) + case models.RelationshipUpdateModeRemove: + return t.destroyJoins(ctx, id, foreignIDs) + } + + return nil +} + +type stashIDTable struct { + table +} + +type stashIDRow struct { + StashID null.String `db:"stash_id"` + Endpoint null.String `db:"endpoint"` + UpdatedAt Timestamp `db:"updated_at"` +} + +func (r *stashIDRow) resolve() models.StashID { + return models.StashID{ + StashID: r.StashID.String, + Endpoint: r.Endpoint.String, + UpdatedAt: r.UpdatedAt.Timestamp.UTC(), + } +} + +func (t *stashIDTable) get(ctx context.Context, id int) ([]models.StashID, error) { + q := dialect.Select("endpoint", "stash_id", "updated_at").From(t.table.table).Where(t.idColumn.Eq(id)) + + const single = false + var ret []models.StashID + if err := queryFunc(ctx, q, single, func(rows *sqlx.Rows) error { + var v stashIDRow + if err := rows.StructScan(&v); err != nil { + return err + } + + ret = append(ret, v.resolve()) + + return nil + }); err != nil { + return nil, fmt.Errorf("getting stash ids from %s: %w", t.table.table.GetTable(), err) + } + + return ret, nil +} + +var epochTime = time.Unix(0, 0).UTC() + +func (t *stashIDTable) insertJoin(ctx context.Context, id int, v models.StashID) (sql.Result, error) { + // #5563 - it's possible that zero-value updated at timestamps are provided via import + // replace them with the epoch time + if v.UpdatedAt.IsZero() { + v.UpdatedAt = epochTime + } + + var q = dialect.Insert(t.table.table).Cols(t.idColumn.GetCol(), "endpoint", "stash_id", "updated_at").Vals( + goqu.Vals{id, v.Endpoint, v.StashID, v.UpdatedAt}, + ) + ret, err := exec(ctx, q) + if err != nil { + return nil, fmt.Errorf("inserting into %s: %w", t.table.table.GetTable(), err) + } + + return ret, nil +} + +func (t *stashIDTable) insertJoins(ctx context.Context, id int, v []models.StashID) error { + for _, fk := range v { + if _, err := t.insertJoin(ctx, id, fk); err != nil { + return err + } + } + + return nil +} + +func (t *stashIDTable) replaceJoins(ctx context.Context, id int, v []models.StashID) error { + if err := t.destroy(ctx, []int{id}); err != nil { + return err + } + + return t.insertJoins(ctx, id, v) +} + +func (t *stashIDTable) addJoins(ctx context.Context, id int, v []models.StashID) error { + // get existing foreign keys + fks, err := t.get(ctx, id) + if err != nil { + return err + } + + // only add values that are not already present + var filtered []models.StashID + for _, vv := range v { + for _, e := range fks { + if vv.Endpoint == e.Endpoint { + continue + } + + filtered = append(filtered, vv) + } + } + return t.insertJoins(ctx, id, filtered) +} + +func (t *stashIDTable) destroyJoins(ctx context.Context, id int, v []models.StashID) error { + for _, vv := range v { + q := dialect.Delete(t.table.table).Where( + t.idColumn.Eq(id), + t.table.table.Col("endpoint").Eq(vv.Endpoint), + t.table.table.Col("stash_id").Eq(vv.StashID), + ) + + if _, err := exec(ctx, q); err != nil { + return fmt.Errorf("destroying %s: %w", t.table.table.GetTable(), err) + } + } + + return nil +} + +func (t *stashIDTable) modifyJoins(ctx context.Context, id int, v []models.StashID, mode models.RelationshipUpdateMode) error { + switch mode { + case models.RelationshipUpdateModeSet: + return t.replaceJoins(ctx, id, v) + case models.RelationshipUpdateModeAdd: + return t.addJoins(ctx, id, v) + case models.RelationshipUpdateModeRemove: + return t.destroyJoins(ctx, id, v) + } + + return nil +} + +type stringTable struct { + table + stringColumn exp.IdentifierExpression +} + +func (t *stringTable) get(ctx context.Context, id int) ([]string, error) { + q := dialect.Select(t.stringColumn).From(t.table.table).Where(t.idColumn.Eq(id)) + + const single = false + var ret []string + if err := queryFunc(ctx, q, single, func(rows *sqlx.Rows) error { + var v string + if err := rows.Scan(&v); err != nil { + return err + } + + ret = append(ret, v) + + return nil + }); err != nil { + return nil, fmt.Errorf("getting stash ids from %s: %w", t.table.table.GetTable(), err) + } + + return ret, nil +} + +func (t *stringTable) insertJoin(ctx context.Context, id int, v string) (sql.Result, error) { + q := dialect.Insert(t.table.table).Cols(t.idColumn.GetCol(), t.stringColumn.GetCol()).Vals( + goqu.Vals{id, v}, + ) + ret, err := exec(ctx, q) + if err != nil { + return nil, fmt.Errorf("inserting into %s: %w", t.table.table.GetTable(), err) + } + + return ret, nil +} + +func (t *stringTable) insertJoins(ctx context.Context, id int, v []string) error { + for _, fk := range v { + if _, err := t.insertJoin(ctx, id, fk); err != nil { + return err + } + } + + return nil +} + +func (t *stringTable) replaceJoins(ctx context.Context, id int, v []string) error { + if err := t.destroy(ctx, []int{id}); err != nil { + return err + } + + return t.insertJoins(ctx, id, v) +} + +func (t *stringTable) addJoins(ctx context.Context, id int, v []string) error { + // get existing foreign keys + existing, err := t.get(ctx, id) + if err != nil { + return err + } + + // only add values that are not already present + filtered := sliceutil.Exclude(v, existing) + return t.insertJoins(ctx, id, filtered) +} + +func (t *stringTable) destroyJoins(ctx context.Context, id int, v []string) error { + for _, vv := range v { + q := dialect.Delete(t.table.table).Where( + t.idColumn.Eq(id), + t.stringColumn.Eq(vv), + ) + + if _, err := exec(ctx, q); err != nil { + return fmt.Errorf("destroying %s: %w", t.table.table.GetTable(), err) + } + } + + return nil +} + +func (t *stringTable) modifyJoins(ctx context.Context, id int, v []string, mode models.RelationshipUpdateMode) error { + switch mode { + case models.RelationshipUpdateModeSet: + return t.replaceJoins(ctx, id, v) + case models.RelationshipUpdateModeAdd: + return t.addJoins(ctx, id, v) + case models.RelationshipUpdateModeRemove: + return t.destroyJoins(ctx, id, v) + } + + return nil +} + +type orderedValueTable[T comparable] struct { + table + valueColumn exp.IdentifierExpression +} + +func (t *orderedValueTable[T]) positionColumn() exp.IdentifierExpression { + const positionColumn = "position" + return t.table.table.Col(positionColumn) +} + +func (t *orderedValueTable[T]) get(ctx context.Context, id int) ([]T, error) { + q := dialect.Select(t.valueColumn).From(t.table.table).Where(t.idColumn.Eq(id)).Order(t.positionColumn().Asc()) + + const single = false + var ret []T + if err := queryFunc(ctx, q, single, func(rows *sqlx.Rows) error { + var v T + if err := rows.Scan(&v); err != nil { + return err + } + + ret = append(ret, v) + + return nil + }); err != nil { + return nil, fmt.Errorf("getting stash ids from %s: %w", t.table.table.GetTable(), err) + } + + return ret, nil +} + +func (t *orderedValueTable[T]) insertJoin(ctx context.Context, id int, position int, v T) (sql.Result, error) { + q := dialect.Insert(t.table.table).Cols(t.idColumn.GetCol(), t.positionColumn().GetCol(), t.valueColumn.GetCol()).Vals( + goqu.Vals{id, position, v}, + ) + ret, err := exec(ctx, q) + if err != nil { + return nil, fmt.Errorf("inserting into %s: %w", t.table.table.GetTable(), err) + } + + return ret, nil +} + +func (t *orderedValueTable[T]) insertJoins(ctx context.Context, id int, startPos int, v []T) error { + for i, fk := range v { + if _, err := t.insertJoin(ctx, id, i+startPos, fk); err != nil { + return err + } + } + + return nil +} + +func (t *orderedValueTable[T]) replaceJoins(ctx context.Context, id int, v []T) error { + if err := t.destroy(ctx, []int{id}); err != nil { + return err + } + + const startPos = 0 + return t.insertJoins(ctx, id, startPos, v) +} + +func (t *orderedValueTable[T]) addJoins(ctx context.Context, id int, v []T) error { + // get existing foreign keys + existing, err := t.get(ctx, id) + if err != nil { + return err + } + + // only add values that are not already present + filtered := sliceutil.Exclude(v, existing) + + if len(filtered) == 0 { + return nil + } + + startPos := len(existing) + return t.insertJoins(ctx, id, startPos, filtered) +} + +func (t *orderedValueTable[T]) destroyJoins(ctx context.Context, id int, v []T) error { + existing, err := t.get(ctx, id) + if err != nil { + return fmt.Errorf("getting existing %s: %w", t.table.table.GetTable(), err) + } + + newValue := sliceutil.Exclude(existing, v) + if len(newValue) == len(existing) { + return nil + } + + return t.replaceJoins(ctx, id, newValue) +} + +func (t *orderedValueTable[T]) modifyJoins(ctx context.Context, id int, v []T, mode models.RelationshipUpdateMode) error { + switch mode { + case models.RelationshipUpdateModeSet: + return t.replaceJoins(ctx, id, v) + case models.RelationshipUpdateModeAdd: + return t.addJoins(ctx, id, v) + case models.RelationshipUpdateModeRemove: + return t.destroyJoins(ctx, id, v) + } + + return nil +} + +type scenesGroupsTable struct { + table +} + +type groupsScenesRow struct { + SceneID null.Int `db:"scene_id"` + GroupID null.Int `db:"group_id"` + SceneIndex null.Int `db:"scene_index"` +} + +func (r groupsScenesRow) resolve(sceneID int) models.GroupsScenes { + return models.GroupsScenes{ + GroupID: int(r.GroupID.Int64), + SceneIndex: nullIntPtr(r.SceneIndex), + } +} + +func (t *scenesGroupsTable) get(ctx context.Context, id int) ([]models.GroupsScenes, error) { + q := dialect.Select("group_id", "scene_index").From(t.table.table).Where(t.idColumn.Eq(id)) + + const single = false + var ret []models.GroupsScenes + if err := queryFunc(ctx, q, single, func(rows *sqlx.Rows) error { + var v groupsScenesRow + if err := rows.StructScan(&v); err != nil { + return err + } + + ret = append(ret, v.resolve(id)) + + return nil + }); err != nil { + return nil, fmt.Errorf("getting scene groups from %s: %w", t.table.table.GetTable(), err) + } + + return ret, nil +} + +func (t *scenesGroupsTable) insertJoin(ctx context.Context, id int, v models.GroupsScenes) (sql.Result, error) { + q := dialect.Insert(t.table.table).Cols(t.idColumn.GetCol(), "group_id", "scene_index").Vals( + goqu.Vals{id, v.GroupID, intFromPtr(v.SceneIndex)}, + ) + ret, err := exec(ctx, q) + if err != nil { + return nil, fmt.Errorf("inserting into %s: %w", t.table.table.GetTable(), err) + } + + return ret, nil +} + +func (t *scenesGroupsTable) insertJoins(ctx context.Context, id int, v []models.GroupsScenes) error { + for _, fk := range v { + if _, err := t.insertJoin(ctx, id, fk); err != nil { + return err + } + } + + return nil +} + +func (t *scenesGroupsTable) replaceJoins(ctx context.Context, id int, v []models.GroupsScenes) error { + if err := t.destroy(ctx, []int{id}); err != nil { + return err + } + + return t.insertJoins(ctx, id, v) +} + +func (t *scenesGroupsTable) addJoins(ctx context.Context, id int, v []models.GroupsScenes) error { + // get existing foreign keys + fks, err := t.get(ctx, id) + if err != nil { + return err + } + + // only add values that are not already present + var filtered []models.GroupsScenes + for _, vv := range v { + found := false + + for _, e := range fks { + if vv.GroupID == e.GroupID { + found = true + break + } + } + + if !found { + filtered = append(filtered, vv) + } + } + return t.insertJoins(ctx, id, filtered) +} + +func (t *scenesGroupsTable) destroyJoins(ctx context.Context, id int, v []models.GroupsScenes) error { + for _, vv := range v { + q := dialect.Delete(t.table.table).Where( + t.idColumn.Eq(id), + t.table.table.Col("group_id").Eq(vv.GroupID), + ) + + if _, err := exec(ctx, q); err != nil { + return fmt.Errorf("destroying %s: %w", t.table.table.GetTable(), err) + } + } + + return nil +} + +func (t *scenesGroupsTable) modifyJoins(ctx context.Context, id int, v []models.GroupsScenes, mode models.RelationshipUpdateMode) error { + switch mode { + case models.RelationshipUpdateModeSet: + return t.replaceJoins(ctx, id, v) + case models.RelationshipUpdateModeAdd: + return t.addJoins(ctx, id, v) + case models.RelationshipUpdateModeRemove: + return t.destroyJoins(ctx, id, v) + } + + return nil +} + +type imageGalleriesTable struct { + joinTable +} + +func (t *imageGalleriesTable) setCover(ctx context.Context, id int, galleryID int) error { + if err := t.resetCover(ctx, galleryID); err != nil { + return err + } + + table := t.table.table + + q := dialect.Update(table).Prepared(true).Set(goqu.Record{ + "cover": true, + }).Where(t.idColumn.Eq(id), table.Col(galleryIDColumn).Eq(galleryID)) + + if _, err := exec(ctx, q); err != nil { + return fmt.Errorf("setting cover flag in %s: %w", t.table.table.GetTable(), err) + } + + return nil +} + +func (t *imageGalleriesTable) resetCover(ctx context.Context, galleryID int) error { + table := t.table.table + + q := dialect.Update(table).Prepared(true).Set(goqu.Record{ + "cover": false, + }).Where( + table.Col(galleryIDColumn).Eq(galleryID), + table.Col("cover").IsTrue(), + ) + + if _, err := exec(ctx, q); err != nil { + return fmt.Errorf("unsetting cover flags in %s: %w", t.table.table.GetTable(), err) + } + + return nil +} + +type relatedFilesTable struct { + table +} + +// type scenesFilesRow struct { +// SceneID int `db:"scene_id"` +// Primary bool `db:"primary"` +// FileID models.FileID `db:"file_id"` +// } + +// get returns the file IDs related to the provided scene ID +// the primary file is returned first +func (t *relatedFilesTable) get(ctx context.Context, id int) ([]models.FileID, error) { + q := dialect.Select("file_id").From(t.table.table).Where(t.idColumn.Eq(id)).Order(t.table.table.Col("primary").Desc()) + + const single = false + var ret []models.FileID + if err := queryFunc(ctx, q, single, func(rows *sqlx.Rows) error { + var v models.FileID + if err := rows.Scan(&v); err != nil { + return err + } + + ret = append(ret, v) + + return nil + }); err != nil { + return nil, fmt.Errorf("getting related files from %s: %w", t.table.table.GetTable(), err) + } + + return ret, nil +} + +func (t *relatedFilesTable) insertJoin(ctx context.Context, id int, primary bool, fileID models.FileID) error { + q := dialect.Insert(t.table.table).Cols(t.idColumn.GetCol(), "primary", "file_id").Vals( + goqu.Vals{id, primary, fileID}, + ) + _, err := exec(ctx, q) + if err != nil { + return fmt.Errorf("inserting into %s: %w", t.table.table.GetTable(), err) + } + + return nil +} + +func (t *relatedFilesTable) insertJoins(ctx context.Context, id int, firstPrimary bool, fileIDs []models.FileID) error { + for i, fk := range fileIDs { + if err := t.insertJoin(ctx, id, firstPrimary && i == 0, fk); err != nil { + return err + } + } + + return nil +} + +func (t *relatedFilesTable) replaceJoins(ctx context.Context, id int, fileIDs []models.FileID) error { + if err := t.destroy(ctx, []int{id}); err != nil { + return err + } + + const firstPrimary = true + return t.insertJoins(ctx, id, firstPrimary, fileIDs) +} + +// destroyJoins destroys all entries in the table with the provided fileIDs +func (t *relatedFilesTable) destroyJoins(ctx context.Context, fileIDs []models.FileID) error { + q := dialect.Delete(t.table.table).Where(t.table.table.Col("file_id").In(fileIDs)) + + if _, err := exec(ctx, q); err != nil { + return fmt.Errorf("destroying file joins in %s: %w", t.table.table.GetTable(), err) + } + + return nil +} + +func (t *relatedFilesTable) setPrimary(ctx context.Context, id int, fileID models.FileID) error { + table := t.table.table + + q := dialect.Update(table).Prepared(true).Set(goqu.Record{ + "primary": false, + }).Where(t.idColumn.Eq(id), table.Col(fileIDColumn).Neq(fileID)) + + if _, err := exec(ctx, q); err != nil { + return fmt.Errorf("unsetting primary flags in %s: %w", t.table.table.GetTable(), err) + } + + q = dialect.Update(table).Prepared(true).Set(goqu.Record{ + "primary": true, + }).Where(t.idColumn.Eq(id), table.Col(fileIDColumn).Eq(fileID)) + + if _, err := exec(ctx, q); err != nil { + return fmt.Errorf("setting primary flag in %s: %w", t.table.table.GetTable(), err) + } + + return nil +} + +type viewHistoryTable struct { + table + dateColumn exp.IdentifierExpression +} + +func (t *viewHistoryTable) getDates(ctx context.Context, id int) ([]time.Time, error) { + table := t.table.table + + q := dialect.Select( + t.dateColumn, + ).From(table).Where( + t.idColumn.Eq(id), + ).Order(t.dateColumn.Desc()) + + const single = false + var ret []time.Time + if err := queryFunc(ctx, q, single, func(rows *sqlx.Rows) error { + var date Timestamp + if err := rows.Scan(&date); err != nil { + return err + } + ret = append(ret, date.Timestamp) + return nil + }); err != nil { + return nil, err + } + + return ret, nil +} + +func (t *viewHistoryTable) getManyDates(ctx context.Context, ids []int) ([][]time.Time, error) { + table := t.table.table + + q := dialect.Select( + t.idColumn, + t.dateColumn, + ).From(table).Where( + t.idColumn.In(ids), + ).Order(t.dateColumn.Desc()) + + ret := make([][]time.Time, len(ids)) + idToIndex := idToIndexMap(ids) + + const single = false + if err := queryFunc(ctx, q, single, func(rows *sqlx.Rows) error { + var id int + var date Timestamp + if err := rows.Scan(&id, &date); err != nil { + return err + } + + idx := idToIndex[id] + ret[idx] = append(ret[idx], date.Timestamp) + return nil + }); err != nil { + return nil, err + } + + return ret, nil +} + +func (t *viewHistoryTable) getLastDate(ctx context.Context, id int) (*time.Time, error) { + table := t.table.table + q := dialect.Select(t.dateColumn).From(table).Where( + t.idColumn.Eq(id), + ).Order(t.dateColumn.Desc()).Limit(1) + + var date NullTimestamp + if err := querySimple(ctx, q, &date); err != nil { + return nil, err + } + + return date.TimePtr(), nil +} + +func (t *viewHistoryTable) getManyLastDate(ctx context.Context, ids []int) ([]*time.Time, error) { + table := t.table.table + + q := dialect.Select( + t.idColumn, + goqu.MAX(t.dateColumn), + ).From(table).Where( + t.idColumn.In(ids), + ).GroupBy(t.idColumn) + + ret := make([]*time.Time, len(ids)) + idToIndex := idToIndexMap(ids) + + const single = false + if err := queryFunc(ctx, q, single, func(rows *sqlx.Rows) error { + var id int + + // MAX appears to return a string, so handle it manually + var dateString string + + if err := rows.Scan(&id, &dateString); err != nil { + return err + } + + t, err := time.Parse(TimestampFormat, dateString) + if err != nil { + return fmt.Errorf("parsing date %v: %w", dateString, err) + } + + idx := idToIndex[id] + ret[idx] = &t + return nil + }); err != nil { + return nil, err + } + + return ret, nil +} + +func (t *viewHistoryTable) getCount(ctx context.Context, id int) (int, error) { + table := t.table.table + q := dialect.Select(goqu.COUNT("*")).From(table).Where(t.idColumn.Eq(id)) + + const single = true + var ret int + if err := queryFunc(ctx, q, single, func(rows *sqlx.Rows) error { + if err := rows.Scan(&ret); err != nil { + return err + } + return nil + }); err != nil { + return 0, err + } + + return ret, nil +} + +func (t *viewHistoryTable) getManyCount(ctx context.Context, ids []int) ([]int, error) { + table := t.table.table + + q := dialect.Select( + t.idColumn, + goqu.COUNT(t.dateColumn), + ).From(table).Where( + t.idColumn.In(ids), + ).GroupBy(t.idColumn) + + ret := make([]int, len(ids)) + idToIndex := idToIndexMap(ids) + + const single = false + if err := queryFunc(ctx, q, single, func(rows *sqlx.Rows) error { + var id int + var count int + if err := rows.Scan(&id, &count); err != nil { + return err + } + + idx := idToIndex[id] + ret[idx] = count + return nil + }); err != nil { + return nil, err + } + + return ret, nil +} + +func (t *viewHistoryTable) getAllCount(ctx context.Context) (int, error) { + table := t.table.table + q := dialect.Select(goqu.COUNT("*")).From(table) + + const single = true + var ret int + if err := queryFunc(ctx, q, single, func(rows *sqlx.Rows) error { + if err := rows.Scan(&ret); err != nil { + return err + } + return nil + }); err != nil { + return 0, err + } + + return ret, nil +} + +func (t *viewHistoryTable) getUniqueCount(ctx context.Context) (int, error) { + table := t.table.table + q := dialect.Select(goqu.COUNT(goqu.DISTINCT(t.idColumn))).From(table) + + const single = true + var ret int + if err := queryFunc(ctx, q, single, func(rows *sqlx.Rows) error { + if err := rows.Scan(&ret); err != nil { + return err + } + return nil + }); err != nil { + return 0, err + } + + return ret, nil +} + +func (t *viewHistoryTable) addDates(ctx context.Context, id int, dates []time.Time) ([]time.Time, error) { + table := t.table.table + + if len(dates) == 0 { + dates = []time.Time{time.Now()} + } + + for _, d := range dates { + q := dialect.Insert(table).Cols(t.idColumn.GetCol(), t.dateColumn.GetCol()).Vals( + // convert all dates to UTC + goqu.Vals{id, UTCTimestamp{Timestamp{d}}}, + ) + + if _, err := exec(ctx, q); err != nil { + return nil, fmt.Errorf("inserting into %s: %w", table.GetTable(), err) + } + } + + return t.getDates(ctx, id) +} + +func (t *viewHistoryTable) deleteDates(ctx context.Context, id int, dates []time.Time) ([]time.Time, error) { + table := t.table.table + + mostRecent := false + if len(dates) == 0 { + mostRecent = true + dates = []time.Time{time.Now()} + } + + for _, date := range dates { + var subquery *goqu.SelectDataset + if mostRecent { + // delete the most recent + subquery = dialect.Select("ctid").From(table).Where( + t.idColumn.Eq(id), + ).Order(t.dateColumn.Desc()).Limit(1) + } else { + subquery = dialect.Select("ctid").From(table).Where( + t.idColumn.Eq(id), + t.dateColumn.Eq(UTCTimestamp{Timestamp{date}}), + ).Limit(1) + } + + q := dialect.Delete(table).Where(goqu.I("ctid").Eq(subquery)) + + if _, err := exec(ctx, q); err != nil { + return nil, fmt.Errorf("deleting from %s: %w", table.GetTable(), err) + } + } + + return t.getDates(ctx, id) +} + +func (t *viewHistoryTable) deleteAllDates(ctx context.Context, id int) (int, error) { + table := t.table.table + q := dialect.Delete(table).Where(t.idColumn.Eq(id)) + + if _, err := exec(ctx, q); err != nil { + return 0, fmt.Errorf("resetting dates for id %v: %w", id, err) + } + + return t.getCount(ctx, id) +} + +type sqler interface { + ToSQL() (sql string, params []interface{}, err error) +} + +func exec(ctx context.Context, stmt sqler) (sql.Result, error) { + tx, err := getTx(ctx) + if err != nil { + return nil, err + } + + sql, args, err := stmt.ToSQL() + if err != nil { + return nil, fmt.Errorf("generating sql: %w", err) + } + + logger.Tracef("SQL: %s [%v]", sql, args) + ret, err := tx.ExecContext(ctx, sql, args...) + if err != nil { + return nil, fmt.Errorf("executing `%s` [%v]: %w", sql, args, err) + } + + return ret, nil +} + +// Execute, but returns an ID +func execID(ctx context.Context, stmt sqler) (*int64, error) { + tx, err := getTx(ctx) + if err != nil { + return nil, err + } + + sql, args, err := stmt.ToSQL() + if err != nil { + return nil, fmt.Errorf("generating sql: %w", err) + } + + logger.Tracef("SQL: %s [%v]", sql, args) + var id int64 + err = tx.QueryRowContext(ctx, sql, args...).Scan(&id) + if err != nil { + return nil, fmt.Errorf("executing `%s` [%v]: %w", sql, args, err) + } + + return &id, nil +} + +func count(ctx context.Context, q *goqu.SelectDataset) (int, error) { + var count int + if err := querySimple(ctx, q, &count); err != nil { + return 0, err + } + + return count, nil +} + +func queryFunc(ctx context.Context, query *goqu.SelectDataset, single bool, f func(rows *sqlx.Rows) error) error { + q, args, err := query.ToSQL() + if err != nil { + return err + } + + rows, err := dbWrapper.QueryxContext(ctx, q, args...) + + if err != nil && !errors.Is(err, sql.ErrNoRows) { + return fmt.Errorf("querying `%s` [%v]: %w", q, args, err) + } + defer rows.Close() + + for rows.Next() { + if err := f(rows); err != nil { + return err + } + if single { + break + } + } + + if err := rows.Err(); err != nil { + return err + } + + return nil +} + +func querySimple(ctx context.Context, query *goqu.SelectDataset, out interface{}) error { + q, args, err := query.ToSQL() + if err != nil { + return err + } + + rows, err := dbWrapper.QueryxContext(ctx, q, args...) + if err != nil { + return fmt.Errorf("querying `%s` [%v]: %w", q, args, err) + } + defer rows.Close() + + if rows.Next() { + if err := rows.Scan(out); err != nil { + return err + } + } + + if err := rows.Err(); err != nil { + return err + } + + return nil +} + +// func cols(table exp.IdentifierExpression, cols []string) []interface{} { +// var ret []interface{} +// for _, c := range cols { +// ret = append(ret, table.Col(c)) +// } +// return ret +// } diff --git a/pkg/postgres/tables.go b/pkg/postgres/tables.go new file mode 100644 index 0000000000..34faa81416 --- /dev/null +++ b/pkg/postgres/tables.go @@ -0,0 +1,429 @@ +package postgres + +import ( + "github.com/doug-martin/goqu/v9" + + _ "github.com/doug-martin/goqu/v9/dialect/postgres" +) + +var dialect = goqu.Dialect("postgres") + +var ( + galleriesImagesJoinTable = goqu.T(galleriesImagesTable) + imagesTagsJoinTable = goqu.T(imagesTagsTable) + performersImagesJoinTable = goqu.T(performersImagesTable) + imagesFilesJoinTable = goqu.T(imagesFilesTable) + imagesURLsJoinTable = goqu.T(imagesURLsTable) + + galleriesFilesJoinTable = goqu.T(galleriesFilesTable) + galleriesTagsJoinTable = goqu.T(galleriesTagsTable) + performersGalleriesJoinTable = goqu.T(performersGalleriesTable) + galleriesScenesJoinTable = goqu.T(galleriesScenesTable) + galleriesURLsJoinTable = goqu.T(galleriesURLsTable) + + scenesFilesJoinTable = goqu.T(scenesFilesTable) + scenesTagsJoinTable = goqu.T(scenesTagsTable) + scenesPerformersJoinTable = goqu.T(performersScenesTable) + scenesStashIDsJoinTable = goqu.T("scene_stash_ids") + scenesGroupsJoinTable = goqu.T(groupsScenesTable) + scenesURLsJoinTable = goqu.T(scenesURLsTable) + + sceneMarkersTagsJoinTable = goqu.T(sceneMarkersTagsTable) + + performersAliasesJoinTable = goqu.T(performersAliasesTable) + performersURLsJoinTable = goqu.T(performerURLsTable) + performersTagsJoinTable = goqu.T(performersTagsTable) + performersStashIDsJoinTable = goqu.T("performer_stash_ids") + performersCustomFieldsTable = goqu.T("performer_custom_fields") + + studiosAliasesJoinTable = goqu.T(studioAliasesTable) + studiosURLsJoinTable = goqu.T(studioURLsTable) + studiosTagsJoinTable = goqu.T(studiosTagsTable) + studiosStashIDsJoinTable = goqu.T("studio_stash_ids") + + groupsURLsJoinTable = goqu.T(groupURLsTable) + groupsTagsJoinTable = goqu.T(groupsTagsTable) + groupRelationsJoinTable = goqu.T(groupRelationsTable) + + tagsAliasesJoinTable = goqu.T(tagAliasesTable) + tagRelationsJoinTable = goqu.T(tagRelationsTable) + tagsStashIDsJoinTable = goqu.T("tag_stash_ids") +) + +var ( + imageTableMgr = &table{ + table: goqu.T(imageTable), + idColumn: goqu.T(imageTable).Col(idColumn), + } + + imagesFilesTableMgr = &relatedFilesTable{ + table: table{ + table: imagesFilesJoinTable, + idColumn: imagesFilesJoinTable.Col(imageIDColumn), + }, + } + + imageGalleriesTableMgr = &imageGalleriesTable{ + joinTable: joinTable{ + table: table{ + table: galleriesImagesJoinTable, + idColumn: galleriesImagesJoinTable.Col(imageIDColumn), + }, + fkColumn: galleriesImagesJoinTable.Col(galleryIDColumn), + }, + } + + imagesTagsTableMgr = &joinTable{ + table: table{ + table: imagesTagsJoinTable, + idColumn: imagesTagsJoinTable.Col(imageIDColumn), + }, + fkColumn: imagesTagsJoinTable.Col(tagIDColumn), + foreignTable: tagTableMgr, + orderBy: tagTableSort, + } + + imagesPerformersTableMgr = &joinTable{ + table: table{ + table: performersImagesJoinTable, + idColumn: performersImagesJoinTable.Col(imageIDColumn), + }, + fkColumn: performersImagesJoinTable.Col(performerIDColumn), + } + + imagesURLsTableMgr = &orderedValueTable[string]{ + table: table{ + table: imagesURLsJoinTable, + idColumn: imagesURLsJoinTable.Col(imageIDColumn), + }, + valueColumn: imagesURLsJoinTable.Col(imageURLColumn), + } +) + +var ( + galleryTableMgr = &table{ + table: goqu.T(galleryTable), + idColumn: goqu.T(galleryTable).Col(idColumn), + } + + galleriesFilesTableMgr = &relatedFilesTable{ + table: table{ + table: galleriesFilesJoinTable, + idColumn: galleriesFilesJoinTable.Col(galleryIDColumn), + }, + } + + galleriesTagsTableMgr = &joinTable{ + table: table{ + table: galleriesTagsJoinTable, + idColumn: galleriesTagsJoinTable.Col(galleryIDColumn), + }, + fkColumn: galleriesTagsJoinTable.Col(tagIDColumn), + foreignTable: tagTableMgr, + orderBy: tagTableSort, + } + + galleriesPerformersTableMgr = &joinTable{ + table: table{ + table: performersGalleriesJoinTable, + idColumn: performersGalleriesJoinTable.Col(galleryIDColumn), + }, + fkColumn: performersGalleriesJoinTable.Col(performerIDColumn), + } + + galleriesScenesTableMgr = &joinTable{ + table: table{ + table: galleriesScenesJoinTable, + idColumn: galleriesScenesJoinTable.Col(galleryIDColumn), + }, + fkColumn: galleriesScenesJoinTable.Col(sceneIDColumn), + } + + galleriesChaptersTableMgr = &table{ + table: goqu.T(galleriesChaptersTable), + idColumn: goqu.T(galleriesChaptersTable).Col(idColumn), + } + + galleriesURLsTableMgr = &orderedValueTable[string]{ + table: table{ + table: galleriesURLsJoinTable, + idColumn: galleriesURLsJoinTable.Col(galleryIDColumn), + }, + valueColumn: galleriesURLsJoinTable.Col(galleriesURLColumn), + } +) + +var ( + sceneTableMgr = &table{ + table: goqu.T(sceneTable), + idColumn: goqu.T(sceneTable).Col(idColumn), + } + + sceneMarkerTableMgr = &table{ + table: goqu.T(sceneMarkerTable), + idColumn: goqu.T(sceneMarkerTable).Col(idColumn), + } + + scenesFilesTableMgr = &relatedFilesTable{ + table: table{ + table: scenesFilesJoinTable, + idColumn: scenesFilesJoinTable.Col(sceneIDColumn), + }, + } + + sceneMarkersTagsTableMgr = &joinTable{ + table: table{ + table: sceneMarkersTagsJoinTable, + idColumn: sceneMarkersTagsJoinTable.Col(sceneMarkerIDColumn), + }, + fkColumn: sceneMarkersTagsJoinTable.Col(tagIDColumn), + foreignTable: tagTableMgr, + orderBy: tagTableSort, + } + + scenesTagsTableMgr = &joinTable{ + table: table{ + table: scenesTagsJoinTable, + idColumn: scenesTagsJoinTable.Col(sceneIDColumn), + }, + fkColumn: scenesTagsJoinTable.Col(tagIDColumn), + foreignTable: tagTableMgr, + orderBy: tagTableSort, + } + + scenesPerformersTableMgr = &joinTable{ + table: table{ + table: scenesPerformersJoinTable, + idColumn: scenesPerformersJoinTable.Col(sceneIDColumn), + }, + fkColumn: scenesPerformersJoinTable.Col(performerIDColumn), + } + + scenesGalleriesTableMgr = galleriesScenesTableMgr.invert() + + scenesStashIDsTableMgr = &stashIDTable{ + table: table{ + table: scenesStashIDsJoinTable, + idColumn: scenesStashIDsJoinTable.Col(sceneIDColumn), + }, + } + + scenesGroupsTableMgr = &scenesGroupsTable{ + table: table{ + table: scenesGroupsJoinTable, + idColumn: scenesGroupsJoinTable.Col(sceneIDColumn), + }, + } + + scenesURLsTableMgr = &orderedValueTable[string]{ + table: table{ + table: scenesURLsJoinTable, + idColumn: scenesURLsJoinTable.Col(sceneIDColumn), + }, + valueColumn: scenesURLsJoinTable.Col(sceneURLColumn), + } + + scenesViewTableMgr = &viewHistoryTable{ + table: table{ + table: goqu.T(scenesViewDatesTable), + idColumn: goqu.T(scenesViewDatesTable).Col(sceneIDColumn), + }, + dateColumn: goqu.T(scenesViewDatesTable).Col(sceneViewDateColumn), + } + + scenesOTableMgr = &viewHistoryTable{ + table: table{ + table: goqu.T(scenesODatesTable), + idColumn: goqu.T(scenesODatesTable).Col(sceneIDColumn), + }, + dateColumn: goqu.T(scenesODatesTable).Col(sceneODateColumn), + } +) + +var ( + fileTableMgr = &table{ + table: goqu.T(fileTable), + idColumn: goqu.T(fileTable).Col(idColumn), + } + + videoFileTableMgr = &table{ + table: goqu.T(videoFileTable), + idColumn: goqu.T(videoFileTable).Col(fileIDColumn), + } + + imageFileTableMgr = &table{ + table: goqu.T(imageFileTable), + idColumn: goqu.T(imageFileTable).Col(fileIDColumn), + } + + folderTableMgr = &table{ + table: goqu.T(folderTable), + idColumn: goqu.T(folderTable).Col(idColumn), + } + + fingerprintTableMgr = &table{ + table: goqu.T(fingerprintTable), + idColumn: goqu.T(fingerprintTable).Col(idColumn), + } +) + +var ( + performerTableMgr = &table{ + table: goqu.T(performerTable), + idColumn: goqu.T(performerTable).Col(idColumn), + } + + performersAliasesTableMgr = &stringTable{ + table: table{ + table: performersAliasesJoinTable, + idColumn: performersAliasesJoinTable.Col(performerIDColumn), + }, + stringColumn: performersAliasesJoinTable.Col(performerAliasColumn), + } + + performersURLsTableMgr = &orderedValueTable[string]{ + table: table{ + table: performersURLsJoinTable, + idColumn: performersURLsJoinTable.Col(performerIDColumn), + }, + valueColumn: performersURLsJoinTable.Col(performerURLColumn), + } + + performersTagsTableMgr = &joinTable{ + table: table{ + table: performersTagsJoinTable, + idColumn: performersTagsJoinTable.Col(performerIDColumn), + }, + fkColumn: performersTagsJoinTable.Col(tagIDColumn), + foreignTable: tagTableMgr, + orderBy: tagTableSort, + } + + performersStashIDsTableMgr = &stashIDTable{ + table: table{ + table: performersStashIDsJoinTable, + idColumn: performersStashIDsJoinTable.Col(performerIDColumn), + }, + } +) + +var ( + studioTableMgr = &table{ + table: goqu.T(studioTable), + idColumn: goqu.T(studioTable).Col(idColumn), + } + + studiosAliasesTableMgr = &stringTable{ + table: table{ + table: studiosAliasesJoinTable, + idColumn: studiosAliasesJoinTable.Col(studioIDColumn), + }, + stringColumn: studiosAliasesJoinTable.Col(studioAliasColumn), + } + + studiosURLsTableMgr = &orderedValueTable[string]{ + table: table{ + table: studiosURLsJoinTable, + idColumn: studiosURLsJoinTable.Col(studioIDColumn), + }, + valueColumn: studiosURLsJoinTable.Col(studioURLColumn), + } + + studiosTagsTableMgr = &joinTable{ + table: table{ + table: studiosTagsJoinTable, + idColumn: studiosTagsJoinTable.Col(studioIDColumn), + }, + fkColumn: studiosTagsJoinTable.Col(tagIDColumn), + foreignTable: tagTableMgr, + orderBy: tagTableSort, + } + + studiosStashIDsTableMgr = &stashIDTable{ + table: table{ + table: studiosStashIDsJoinTable, + idColumn: studiosStashIDsJoinTable.Col(studioIDColumn), + }, + } +) + +var ( + tagTableMgr = &table{ + table: goqu.T(tagTable), + idColumn: goqu.T(tagTable).Col(idColumn), + } + + // formerly: goqu.COALESCE(tagTableMgr.table.Col("sort_name"), tagTableMgr.table.Col("name")).Asc() + tagTableSort = goqu.L("COALESCE(tags.sort_name, tags.name) COLLATE NATURAL_CI").Asc() + tagTableSortSQL = "COALESCE(tags.sort_name, tags.name) COLLATE NATURAL_CI ASC" + + tagsAliasesTableMgr = &stringTable{ + table: table{ + table: tagsAliasesJoinTable, + idColumn: tagsAliasesJoinTable.Col(tagIDColumn), + }, + stringColumn: tagsAliasesJoinTable.Col(tagAliasColumn), + } + + tagsParentTagsTableMgr = &joinTable{ + table: table{ + table: tagRelationsJoinTable, + idColumn: tagRelationsJoinTable.Col(tagChildIDColumn), + }, + fkColumn: tagRelationsJoinTable.Col(tagParentIDColumn), + foreignTable: tagTableMgr, + orderBy: tagTableSort, + } + + tagsChildTagsTableMgr = *tagsParentTagsTableMgr.invert() + + tagsStashIDsTableMgr = &stashIDTable{ + table: table{ + table: tagsStashIDsJoinTable, + idColumn: tagsStashIDsJoinTable.Col(tagIDColumn), + }, + } +) + +var ( + groupTableMgr = &table{ + table: goqu.T(groupTable), + idColumn: goqu.T(groupTable).Col(idColumn), + } + + groupsURLsTableMgr = &orderedValueTable[string]{ + table: table{ + table: groupsURLsJoinTable, + idColumn: groupsURLsJoinTable.Col(groupIDColumn), + }, + valueColumn: groupsURLsJoinTable.Col(groupURLColumn), + } + + groupsTagsTableMgr = &joinTable{ + table: table{ + table: groupsTagsJoinTable, + idColumn: groupsTagsJoinTable.Col(groupIDColumn), + }, + fkColumn: groupsTagsJoinTable.Col(tagIDColumn), + foreignTable: tagTableMgr, + orderBy: tagTableSort, + } + + groupRelationshipTableMgr = &table{ + table: groupRelationsJoinTable, + } +) + +var ( + blobTableMgr = &table{ + table: goqu.T(blobTable), + idColumn: goqu.T(blobTable).Col(blobChecksumColumn), + } +) + +var ( + savedFilterTableMgr = &table{ + table: goqu.T(savedFilterTable), + idColumn: goqu.T(savedFilterTable).Col(idColumn), + } +) diff --git a/pkg/postgres/tag.go b/pkg/postgres/tag.go new file mode 100644 index 0000000000..22550ed672 --- /dev/null +++ b/pkg/postgres/tag.go @@ -0,0 +1,1041 @@ +package postgres + +import ( + "context" + "database/sql" + "errors" + "fmt" + "slices" + "strings" + + "github.com/doug-martin/goqu/v9" + "github.com/doug-martin/goqu/v9/exp" + "github.com/jmoiron/sqlx" + "gopkg.in/guregu/null.v4" + "gopkg.in/guregu/null.v4/zero" + + "github.com/stashapp/stash/pkg/models" +) + +const ( + tagTable = "tags" + tagIDColumn = "tag_id" + tagAliasesTable = "tag_aliases" + tagAliasColumn = "alias" + + tagImageBlobColumn = "image_blob" + + tagRelationsTable = "tags_relations" + tagParentIDColumn = "parent_id" + tagChildIDColumn = "child_id" +) + +type tagRow struct { + ID int `db:"id" goqu:"skipinsert"` + Name null.String `db:"name"` // TODO: make schema non-nullable + SortName zero.String `db:"sort_name"` + Favorite bool `db:"favorite"` + Description zero.String `db:"description"` + IgnoreAutoTag bool `db:"ignore_auto_tag"` + CreatedAt Timestamp `db:"created_at"` + UpdatedAt Timestamp `db:"updated_at"` + + // not used in resolutions or updates + ImageBlob zero.String `db:"image_blob"` +} + +func (r *tagRow) fromTag(o models.Tag) { + r.ID = o.ID + r.Name = null.StringFrom(o.Name) + r.SortName = zero.StringFrom((o.SortName)) + r.Favorite = o.Favorite + r.Description = zero.StringFrom(o.Description) + r.IgnoreAutoTag = o.IgnoreAutoTag + r.CreatedAt = Timestamp{Timestamp: o.CreatedAt} + r.UpdatedAt = Timestamp{Timestamp: o.UpdatedAt} +} + +func (r *tagRow) resolve() *models.Tag { + ret := &models.Tag{ + ID: r.ID, + Name: r.Name.String, + SortName: r.SortName.String, + Favorite: r.Favorite, + Description: r.Description.String, + IgnoreAutoTag: r.IgnoreAutoTag, + CreatedAt: r.CreatedAt.Timestamp.UTC(), + UpdatedAt: r.UpdatedAt.Timestamp.UTC(), + } + + return ret +} + +type tagPathRow struct { + tagRow + Path string `db:"path"` +} + +func (r *tagPathRow) resolve() *models.TagPath { + ret := &models.TagPath{ + Tag: *r.tagRow.resolve(), + Path: r.Path, + } + + return ret +} + +type tagRowRecord struct { + updateRecord +} + +func (r *tagRowRecord) fromPartial(o models.TagPartial) { + r.setString("name", o.Name) + r.setNullString("sort_name", o.SortName) + r.setNullString("description", o.Description) + r.setBool("favorite", o.Favorite) + r.setBool("ignore_auto_tag", o.IgnoreAutoTag) + r.setTimestamp("created_at", o.CreatedAt) + r.setTimestamp("updated_at", o.UpdatedAt) +} + +type tagRepositoryType struct { + repository + + aliases stringRepository + stashIDs stashIDRepository + + scenes joinRepository + images joinRepository + galleries joinRepository +} + +var ( + tagRepository = tagRepositoryType{ + repository: repository{ + tableName: tagTable, + idColumn: idColumn, + }, + aliases: stringRepository{ + repository: repository{ + tableName: tagAliasesTable, + idColumn: tagIDColumn, + }, + stringColumn: tagAliasColumn, + }, + stashIDs: stashIDRepository{ + repository{ + tableName: "tag_stash_ids", + idColumn: tagIDColumn, + }, + }, + scenes: joinRepository{ + repository: repository{ + tableName: scenesTagsTable, + idColumn: tagIDColumn, + }, + fkColumn: sceneIDColumn, + foreignTable: sceneTable, + }, + images: joinRepository{ + repository: repository{ + tableName: imagesTagsTable, + idColumn: tagIDColumn, + }, + fkColumn: imageIDColumn, + foreignTable: imageTable, + }, + galleries: joinRepository{ + repository: repository{ + tableName: galleriesTagsTable, + idColumn: tagIDColumn, + }, + fkColumn: galleryIDColumn, + foreignTable: galleryTable, + }, + } +) + +type TagStore struct { + blobJoinQueryBuilder + + tableMgr *table +} + +func NewTagStore(blobStore *BlobStore) *TagStore { + return &TagStore{ + blobJoinQueryBuilder: blobJoinQueryBuilder{ + blobStore: blobStore, + joinTable: tagTable, + }, + tableMgr: tagTableMgr, + } +} + +func (qb *TagStore) table() exp.IdentifierExpression { + return qb.tableMgr.table +} + +func (qb *TagStore) selectDataset() *goqu.SelectDataset { + return dialect.From(qb.table()).Select(qb.table().All()) +} + +func (qb *TagStore) Create(ctx context.Context, newObject *models.Tag) error { + var r tagRow + r.fromTag(*newObject) + + id, err := qb.tableMgr.insertID(ctx, r) + if err != nil { + return err + } + + if newObject.Aliases.Loaded() { + if err := tagsAliasesTableMgr.insertJoins(ctx, id, newObject.Aliases.List()); err != nil { + return err + } + } + + if newObject.ParentIDs.Loaded() { + if err := tagsParentTagsTableMgr.insertJoins(ctx, id, newObject.ParentIDs.List()); err != nil { + return err + } + } + + if newObject.ChildIDs.Loaded() { + if err := tagsChildTagsTableMgr.insertJoins(ctx, id, newObject.ChildIDs.List()); err != nil { + return err + } + } + + if newObject.StashIDs.Loaded() { + if err := tagsStashIDsTableMgr.insertJoins(ctx, id, newObject.StashIDs.List()); err != nil { + return err + } + } + + updated, err := qb.find(ctx, id) + if err != nil { + return fmt.Errorf("finding after create: %w", err) + } + + *newObject = *updated + + return nil +} + +func (qb *TagStore) UpdatePartial(ctx context.Context, id int, partial models.TagPartial) (*models.Tag, error) { + r := tagRowRecord{ + updateRecord{ + Record: make(exp.Record), + }, + } + + r.fromPartial(partial) + + if len(r.Record) > 0 { + if err := qb.tableMgr.updateByID(ctx, id, r.Record); err != nil { + return nil, err + } + } + + if partial.Aliases != nil { + if err := tagsAliasesTableMgr.modifyJoins(ctx, id, partial.Aliases.Values, partial.Aliases.Mode); err != nil { + return nil, err + } + } + + if partial.ParentIDs != nil { + if err := tagsParentTagsTableMgr.modifyJoins(ctx, id, partial.ParentIDs.IDs, partial.ParentIDs.Mode); err != nil { + return nil, err + } + } + + if partial.ChildIDs != nil { + if err := tagsChildTagsTableMgr.modifyJoins(ctx, id, partial.ChildIDs.IDs, partial.ChildIDs.Mode); err != nil { + return nil, err + } + } + + if partial.StashIDs != nil { + if err := tagsStashIDsTableMgr.modifyJoins(ctx, id, partial.StashIDs.StashIDs, partial.StashIDs.Mode); err != nil { + return nil, err + } + } + + return qb.find(ctx, id) +} + +func (qb *TagStore) Update(ctx context.Context, updatedObject *models.Tag) error { + var r tagRow + r.fromTag(*updatedObject) + + if err := qb.tableMgr.updateByID(ctx, updatedObject.ID, r); err != nil { + return err + } + + if updatedObject.Aliases.Loaded() { + if err := tagsAliasesTableMgr.replaceJoins(ctx, updatedObject.ID, updatedObject.Aliases.List()); err != nil { + return err + } + } + + if updatedObject.ParentIDs.Loaded() { + if err := tagsParentTagsTableMgr.replaceJoins(ctx, updatedObject.ID, updatedObject.ParentIDs.List()); err != nil { + return err + } + } + + if updatedObject.ChildIDs.Loaded() { + if err := tagsChildTagsTableMgr.replaceJoins(ctx, updatedObject.ID, updatedObject.ChildIDs.List()); err != nil { + return err + } + } + + if updatedObject.StashIDs.Loaded() { + if err := tagsStashIDsTableMgr.replaceJoins(ctx, updatedObject.ID, updatedObject.StashIDs.List()); err != nil { + return err + } + } + + return nil +} + +func (qb *TagStore) Destroy(ctx context.Context, id int) error { + // must handle image checksums manually + if err := qb.destroyImage(ctx, id); err != nil { + return err + } + + // cannot unset primary_tag_id in scene_markers because it is not nullable + countQuery := "SELECT COUNT(*) as count FROM scene_markers where primary_tag_id = ?" + args := []interface{}{id} + primaryMarkers, err := tagRepository.runCountQuery(ctx, countQuery, args) + if err != nil { + return err + } + + if primaryMarkers > 0 { + return errors.New("cannot delete tag used as a primary tag in scene markers") + } + + return tagRepository.destroyExisting(ctx, []int{id}) +} + +// returns nil, nil if not found +func (qb *TagStore) Find(ctx context.Context, id int) (*models.Tag, error) { + ret, err := qb.find(ctx, id) + if errors.Is(err, sql.ErrNoRows) { + return nil, nil + } + return ret, err +} + +func (qb *TagStore) FindMany(ctx context.Context, ids []int) ([]*models.Tag, error) { + ret := make([]*models.Tag, len(ids)) + + if len(ids) == 0 { + return ret, nil + } + + table := qb.table() + if err := batchExec(ids, defaultBatchSize, func(batch []int) error { + q := qb.selectDataset().Prepared(true).Where(table.Col(idColumn).In(batch)) + unsorted, err := qb.getMany(ctx, q) + if err != nil { + return err + } + + for _, s := range unsorted { + i := slices.Index(ids, s.ID) + ret[i] = s + } + + return nil + }); err != nil { + return nil, err + } + + for i := range ret { + if ret[i] == nil { + return nil, fmt.Errorf("tag with id %d not found", ids[i]) + } + } + + return ret, nil +} + +// returns nil, sql.ErrNoRows if not found +func (qb *TagStore) find(ctx context.Context, id int) (*models.Tag, error) { + q := qb.selectDataset().Where(qb.tableMgr.byID(id)) + + ret, err := qb.get(ctx, q) + if err != nil { + return nil, err + } + + return ret, nil +} + +// returns nil, sql.ErrNoRows if not found +func (qb *TagStore) get(ctx context.Context, q *goqu.SelectDataset) (*models.Tag, error) { + ret, err := qb.getMany(ctx, q) + if err != nil { + return nil, err + } + + if len(ret) == 0 { + return nil, sql.ErrNoRows + } + + return ret[0], nil +} + +func (qb *TagStore) getMany(ctx context.Context, q *goqu.SelectDataset) ([]*models.Tag, error) { + const single = false + var ret []*models.Tag + if err := queryFunc(ctx, q, single, func(r *sqlx.Rows) error { + var f tagRow + if err := r.StructScan(&f); err != nil { + return err + } + + s := f.resolve() + + ret = append(ret, s) + return nil + }); err != nil { + return nil, err + } + + return ret, nil +} + +func (qb *TagStore) FindBySceneID(ctx context.Context, sceneID int) ([]*models.Tag, error) { + query := ` + SELECT tags.* FROM tags + LEFT JOIN scenes_tags as scenes_join on scenes_join.tag_id = tags.id + WHERE scenes_join.scene_id = ? + GROUP BY tags.id + ` + add, _ := qb.getDefaultTagSort() + query += add + args := []interface{}{sceneID} + return qb.queryTags(ctx, query, args) +} + +func (qb *TagStore) FindByPerformerID(ctx context.Context, performerID int) ([]*models.Tag, error) { + query := ` + SELECT tags.* FROM tags + LEFT JOIN performers_tags as performers_join on performers_join.tag_id = tags.id + WHERE performers_join.performer_id = ? + GROUP BY tags.id + ` + add, _ := qb.getDefaultTagSort() + query += add + args := []interface{}{performerID} + return qb.queryTags(ctx, query, args) +} + +func (qb *TagStore) FindByImageID(ctx context.Context, imageID int) ([]*models.Tag, error) { + query := ` + SELECT tags.* FROM tags + LEFT JOIN images_tags as images_join on images_join.tag_id = tags.id + WHERE images_join.image_id = ? + GROUP BY tags.id + ` + add, _ := qb.getDefaultTagSort() + query += add + args := []interface{}{imageID} + return qb.queryTags(ctx, query, args) +} + +func (qb *TagStore) FindByGalleryID(ctx context.Context, galleryID int) ([]*models.Tag, error) { + query := ` + SELECT tags.* FROM tags + LEFT JOIN galleries_tags as galleries_join on galleries_join.tag_id = tags.id + WHERE galleries_join.gallery_id = ? + GROUP BY tags.id + ` + add, _ := qb.getDefaultTagSort() + query += add + args := []interface{}{galleryID} + return qb.queryTags(ctx, query, args) +} + +func (qb *TagStore) FindByGroupID(ctx context.Context, groupID int) ([]*models.Tag, error) { + query := ` + SELECT tags.* FROM tags + LEFT JOIN groups_tags as groups_join on groups_join.tag_id = tags.id + WHERE groups_join.group_id = ? + GROUP BY tags.id + ` + add, _ := qb.getDefaultTagSort() + query += add + args := []interface{}{groupID} + return qb.queryTags(ctx, query, args) +} + +func (qb *TagStore) FindBySceneMarkerID(ctx context.Context, sceneMarkerID int) ([]*models.Tag, error) { + query := ` + SELECT tags.* FROM tags + LEFT JOIN scene_markers_tags as scene_markers_join on scene_markers_join.tag_id = tags.id + WHERE scene_markers_join.scene_marker_id = ? + GROUP BY tags.id + ` + add, _ := qb.getDefaultTagSort() + query += add + args := []interface{}{sceneMarkerID} + return qb.queryTags(ctx, query, args) +} + +func (qb *TagStore) FindByStudioID(ctx context.Context, studioID int) ([]*models.Tag, error) { + query := ` + SELECT tags.* FROM tags + LEFT JOIN studios_tags as studios_join on studios_join.tag_id = tags.id + WHERE studios_join.studio_id = ? + GROUP BY tags.id + ` + add, _ := qb.getDefaultTagSort() + query += add + args := []interface{}{studioID} + return qb.queryTags(ctx, query, args) +} + +func (qb *TagStore) FindByName(ctx context.Context, name string, nocase bool) (*models.Tag, error) { + // query := "SELECT * FROM tags WHERE name = ?" + // if nocase { + // query += " COLLATE NOCASE" + // } + // query += " LIMIT 1" + where := "name = ?" + if nocase { + where += " COLLATE NOCASE" + } + sq := qb.selectDataset().Prepared(true).Where(goqu.L(where, name)).Limit(1) + ret, err := qb.get(ctx, sq) + + if err != nil && !errors.Is(err, sql.ErrNoRows) { + return nil, err + } + + return ret, nil +} + +func (qb *TagStore) FindByNames(ctx context.Context, names []string, nocase bool) ([]*models.Tag, error) { + // query := "SELECT * FROM tags WHERE name" + // if nocase { + // query += " COLLATE NOCASE" + // } + // query += " IN " + getInBinding(len(names)) + where := "name" + if nocase { + where += " COLLATE NOCASE" + } + where += " IN " + getInBinding(len(names)) + var args []interface{} + for _, name := range names { + args = append(args, name) + } + sq := qb.selectDataset().Prepared(true).Where(goqu.L(where, args...)) + ret, err := qb.getMany(ctx, sq) + + if err != nil { + return nil, err + } + + return ret, nil +} + +func (qb *TagStore) FindByStashID(ctx context.Context, stashID models.StashID) ([]*models.Tag, error) { + sq := dialect.From(tagsStashIDsJoinTable).Select(tagsStashIDsJoinTable.Col(tagIDColumn)).Where( + tagsStashIDsJoinTable.Col("stash_id").Eq(stashID.StashID), + tagsStashIDsJoinTable.Col("endpoint").Eq(stashID.Endpoint), + ) + + idsQuery := qb.selectDataset().Where( + qb.table().Col(idColumn).In(sq), + ) + + ret, err := qb.getMany(ctx, idsQuery) + if err != nil { + return nil, fmt.Errorf("getting tags for stash ID %s: %w", stashID.StashID, err) + } + + return ret, nil +} + +func (qb *TagStore) GetParentIDs(ctx context.Context, relatedID int) ([]int, error) { + return tagsParentTagsTableMgr.get(ctx, relatedID) +} + +func (qb *TagStore) GetChildIDs(ctx context.Context, relatedID int) ([]int, error) { + return tagsChildTagsTableMgr.get(ctx, relatedID) +} + +func (qb *TagStore) FindByParentTagID(ctx context.Context, parentID int) ([]*models.Tag, error) { + query := ` + SELECT tags.* FROM tags + INNER JOIN tags_relations ON tags_relations.child_id = tags.id + WHERE tags_relations.parent_id = ? + ` + add, _ := qb.getDefaultTagSort() + query += add + args := []interface{}{parentID} + return qb.queryTags(ctx, query, args) +} + +func (qb *TagStore) FindByChildTagID(ctx context.Context, parentID int) ([]*models.Tag, error) { + query := ` + SELECT tags.* FROM tags + INNER JOIN tags_relations ON tags_relations.parent_id = tags.id + WHERE tags_relations.child_id = ? + ` + add, _ := qb.getDefaultTagSort() + query += add + args := []interface{}{parentID} + return qb.queryTags(ctx, query, args) +} + +func (qb *TagStore) CountByParentTagID(ctx context.Context, parentID int) (int, error) { + q := dialect.Select(goqu.COUNT("*")).From(goqu.T("tags")). + InnerJoin(goqu.T("tags_relations"), goqu.On(goqu.I("tags_relations.parent_id").Eq(goqu.I("tags.id")))). + Where(goqu.I("tags_relations.child_id").Eq(goqu.V(parentID))) // Pass the parentID here + return count(ctx, q) +} + +func (qb *TagStore) CountByChildTagID(ctx context.Context, childID int) (int, error) { + q := dialect.Select(goqu.COUNT("*")).From(goqu.T("tags")). + InnerJoin(goqu.T("tags_relations"), goqu.On(goqu.I("tags_relations.child_id").Eq(goqu.I("tags.id")))). + Where(goqu.I("tags_relations.parent_id").Eq(goqu.V(childID))) // Pass the childID here + return count(ctx, q) +} + +func (qb *TagStore) Count(ctx context.Context) (int, error) { + q := dialect.Select(goqu.COUNT("*")).From(qb.table()) + return count(ctx, q) +} + +func (qb *TagStore) All(ctx context.Context) ([]*models.Tag, error) { + table := qb.table() + + return qb.getMany(ctx, qb.selectDataset().Order( + goqu.L("COALESCE(tags.sort_name, tags.name) COLLATE NATURAL_CI").Asc(), + table.Col(idColumn).Asc(), + )) +} + +func (qb *TagStore) QueryForAutoTag(ctx context.Context, words []string) ([]*models.Tag, error) { + // TODO - Query needs to be changed to support queries of this type, and + // this method should be removed + query := selectAll(tagTable) + query += " LEFT JOIN tag_aliases ON tag_aliases.tag_id = tags.id" + + var whereClauses []string + var args []interface{} + + for _, w := range words { + ww := w + "%" + whereClauses = append(whereClauses, "tags.name ILIKE ?") + args = append(args, ww) + + // include aliases + whereClauses = append(whereClauses, "tag_aliases.alias ILIKE ?") + args = append(args, ww) + } + + whereOr := "(" + strings.Join(whereClauses, " OR ") + ")" + where := strings.Join([]string{ + "tags.ignore_auto_tag = false", + whereOr, + }, " AND ") + return qb.queryTags(ctx, query+" WHERE "+where, args) +} + +func (qb *TagStore) Query(ctx context.Context, tagFilter *models.TagFilterType, findFilter *models.FindFilterType) ([]*models.Tag, int, error) { + if tagFilter == nil { + tagFilter = &models.TagFilterType{} + } + if findFilter == nil { + findFilter = &models.FindFilterType{} + } + + query := tagRepository.newQuery() + distinctIDs(&query, tagTable) + + if q := findFilter.Q; q != nil && *q != "" { + query.join(tagAliasesTable, "", "tag_aliases.tag_id = tags.id") + searchColumns := []string{"tags.name", "tag_aliases.alias", "tags.sort_name"} + query.parseQueryString(searchColumns, *q) + } + + filter := filterBuilderFromHandler(ctx, &tagFilterHandler{ + tagFilter: tagFilter, + }) + + if err := query.addFilter(filter); err != nil { + return nil, 0, err + } + + var err error + var group []string + query.sort, group, err = qb.getTagSort(&query, findFilter) + if err != nil { + return nil, 0, err + } + query.pagination = getPagination(findFilter) + query.addGroupBy(group...) + idsResult, countResult, err := query.executeFind(ctx) + if err != nil { + return nil, 0, err + } + + tags, err := qb.FindMany(ctx, idsResult) + if err != nil { + return nil, 0, err + } + + return tags, countResult, nil +} + +var tagSortOptions = sortOptions{ + "created_at", + "galleries_count", + "groups_count", + "id", + "images_count", + "movies_count", + "studios_count", + "name", + "performers_count", + "random", + "scene_markers_count", + "scenes_count", + "scenes_duration", + "updated_at", +} + +func (qb *TagStore) sortByScenesDuration(direction string) string { + return fmt.Sprintf(` ORDER BY ( + SELECT COALESCE(SUM(video_files.duration), 0) + FROM %s + LEFT JOIN %s ON %s.id = %s.%s + LEFT JOIN %s ON %s.%s = %s.id + LEFT JOIN video_files ON video_files.file_id = %s.file_id + WHERE %s.%s = %s.id + ) %s`, scenesTagsTable, sceneTable, sceneTable, scenesTagsTable, sceneIDColumn, scenesFilesTable, scenesFilesTable, sceneIDColumn, sceneTable, scenesFilesTable, scenesTagsTable, tagIDColumn, tagTable, getSortDirection(direction)) +} + +func (qb *TagStore) getDefaultTagSort() (string, []string) { + return getSort("name", "ASC", "tags") +} + +func (qb *TagStore) getTagSort(query *queryBuilder, findFilter *models.FindFilterType) (string, []string, error) { + var sort string + var direction string + if findFilter == nil { + sort = "name" + direction = "ASC" + } else { + sort = findFilter.GetSort("name") + direction = findFilter.GetDirection() + } + + // CVE-2024-32231 - ensure sort is in the list of allowed sorts + if err := tagSortOptions.validateSort(sort); err != nil { + return "", nil, err + } + + group := []string{} + sortQuery := "" + switch sort { + case "name": + sortQuery += fmt.Sprintf(" ORDER BY COALESCE(tags.sort_name, tags.name) COLLATE NATURAL_CI %s", getSortDirection(direction)) + case "scenes_count": + sortQuery += getCountSort(tagTable, scenesTagsTable, tagIDColumn, direction) + case "scenes_duration": + sortQuery += qb.sortByScenesDuration(direction) + case "scene_markers_count": + sortQuery += fmt.Sprintf(" ORDER BY (SELECT COUNT(*) FROM scene_markers_tags WHERE tags.id = scene_markers_tags.tag_id)+(SELECT COUNT(*) FROM scene_markers WHERE tags.id = scene_markers.primary_tag_id) %s", getSortDirection(direction)) + case "images_count": + sortQuery += getCountSort(tagTable, imagesTagsTable, tagIDColumn, direction) + case "galleries_count": + sortQuery += getCountSort(tagTable, galleriesTagsTable, tagIDColumn, direction) + case "performers_count": + sortQuery += getCountSort(tagTable, performersTagsTable, tagIDColumn, direction) + case "studios_count": + sortQuery += getCountSort(tagTable, studiosTagsTable, tagIDColumn, direction) + case "movies_count", "groups_count": + sortQuery += getCountSort(tagTable, groupsTagsTable, tagIDColumn, direction) + default: + var add string + add, group = getSort(sort, direction, "tags") + sortQuery += add + } + + // Whatever the sorting, always use sort_name/name/id as a final sort + sortQuery += ", COALESCE(tags.name, CAST(tags.id as text)) COLLATE NATURAL_CI ASC" + group = append(group, "tags.name", "tags.id") + return sortQuery, group, nil +} + +func (qb *TagStore) queryTags(ctx context.Context, query string, args []interface{}) ([]*models.Tag, error) { + const single = false + var ret []*models.Tag + if err := tagRepository.queryFunc(ctx, query, args, single, func(r *sqlx.Rows) error { + var f tagRow + if err := r.StructScan(&f); err != nil { + return err + } + + s := f.resolve() + + ret = append(ret, s) + return nil + }); err != nil { + return nil, err + } + + return ret, nil +} + +func (qb *TagStore) queryTagPaths(ctx context.Context, query string, args []interface{}) ([]*models.TagPath, error) { + const single = false + var ret []*models.TagPath + if err := tagRepository.queryFunc(ctx, query, args, single, func(r *sqlx.Rows) error { + var f tagPathRow + if err := r.StructScan(&f); err != nil { + return err + } + + t := f.resolve() + + ret = append(ret, t) + return nil + }); err != nil { + return nil, err + } + + return ret, nil +} + +func (qb *TagStore) GetImage(ctx context.Context, tagID int) ([]byte, error) { + return qb.blobJoinQueryBuilder.GetImage(ctx, tagID, tagImageBlobColumn) +} + +func (qb *TagStore) HasImage(ctx context.Context, tagID int) (bool, error) { + return qb.blobJoinQueryBuilder.HasImage(ctx, tagID, tagImageBlobColumn) +} + +func (qb *TagStore) UpdateImage(ctx context.Context, tagID int, image []byte) error { + return qb.blobJoinQueryBuilder.UpdateImage(ctx, tagID, tagImageBlobColumn, image) +} + +func (qb *TagStore) destroyImage(ctx context.Context, tagID int) error { + return qb.blobJoinQueryBuilder.DestroyImage(ctx, tagID, tagImageBlobColumn) +} + +func (qb *TagStore) GetAliases(ctx context.Context, tagID int) ([]string, error) { + return tagRepository.aliases.get(ctx, tagID) +} + +func (qb *TagStore) UpdateAliases(ctx context.Context, tagID int, aliases []string) error { + return tagRepository.aliases.replace(ctx, tagID, aliases) +} + +func (qb *TagStore) GetStashIDs(ctx context.Context, tagID int) ([]models.StashID, error) { + return tagsStashIDsTableMgr.get(ctx, tagID) +} + +func (qb *TagStore) UpdateStashIDs(ctx context.Context, tagID int, stashIDs []models.StashID) error { + return tagsStashIDsTableMgr.replaceJoins(ctx, tagID, stashIDs) +} + +func (qb *TagStore) Merge(ctx context.Context, source []int, destination int) error { + if len(source) == 0 { + return nil + } + + inBinding := getInBinding(len(source)) + + args := []interface{}{destination} + srcArgs := make([]interface{}, len(source)) + for i, id := range source { + if id == destination { + return errors.New("cannot merge where source == destination") + } + srcArgs[i] = id + } + + args = append(args, srcArgs...) + + tagTables := map[string]string{ + scenesTagsTable: sceneIDColumn, + "scene_markers_tags": "scene_marker_id", + galleriesTagsTable: galleryIDColumn, + imagesTagsTable: imageIDColumn, + "performers_tags": "performer_id", + "studios_tags": "studio_id", + groupsTagsTable: "group_id", + } + + // for each table, update source tag ids to destination tag id, ignoring duplicates + for table, idColumn := range tagTables { + for _, to_migrate_id := range srcArgs { + err := withSavepoint(ctx, func(ctx context.Context) error { + _, err := dbWrapper.Exec(ctx, `UPDATE `+table+` + SET tag_id = $1 + WHERE tag_id = $2 + AND NOT EXISTS(SELECT 1 FROM `+table+` o WHERE o.`+idColumn+` = `+table+`.`+idColumn+` AND o.tag_id = $1)`, + destination, to_migrate_id, + ) + return err + }) + + if err != nil && !qb.repository.isConstraintError(err) { + return err + } + } + + // delete source tag ids from the table where they couldn't be set + if _, err := dbWrapper.Exec(ctx, `DELETE FROM `+table+` WHERE tag_id IN `+inBinding, srcArgs...); err != nil { + return err + } + } + + _, err := dbWrapper.Exec(ctx, "UPDATE "+sceneMarkerTable+" SET primary_tag_id = ? WHERE primary_tag_id IN "+inBinding, args...) + if err != nil { + return err + } + + _, err = dbWrapper.Exec(ctx, "INSERT INTO "+tagAliasesTable+" (tag_id, alias) SELECT ?, name FROM "+tagTable+" WHERE id IN "+inBinding, args...) + if err != nil { + return err + } + + _, err = dbWrapper.Exec(ctx, "UPDATE "+tagAliasesTable+" SET tag_id = ? WHERE tag_id IN "+inBinding, args...) + if err != nil { + return err + } + + // Merge StashIDs - insert non-conflicting source StashIDs as new rows with destination tag_id + _, err = dbWrapper.Exec(ctx, `INSERT INTO tag_stash_ids (tag_id, endpoint, stash_id, updated_at) + SELECT ?, endpoint, stash_id, updated_at FROM tag_stash_ids WHERE tag_id IN `+inBinding+` ON CONFLICT DO NOTHING`, args...) + if err != nil { + return err + } + + // Delete all source StashIDs (including those that conflicted and weren't inserted) + if _, err := dbWrapper.Exec(ctx, `DELETE FROM tag_stash_ids WHERE tag_id IN `+inBinding, srcArgs...); err != nil { + return err + } + + for _, id := range source { + err = qb.Destroy(ctx, id) + if err != nil { + return err + } + } + + return nil +} + +func (qb *TagStore) UpdateParentTags(ctx context.Context, tagID int, parentIDs []int) error { + if _, err := dbWrapper.Exec(ctx, "DELETE FROM tags_relations WHERE child_id = ?", tagID); err != nil { + return err + } + + if len(parentIDs) > 0 { + var args []interface{} + var values []string + for _, parentID := range parentIDs { + values = append(values, "(? , ?)") + args = append(args, parentID, tagID) + } + + query := "INSERT INTO tags_relations (parent_id, child_id) VALUES " + strings.Join(values, ", ") + if _, err := dbWrapper.Exec(ctx, query, args...); err != nil { + return err + } + } + + return nil +} + +func (qb *TagStore) UpdateChildTags(ctx context.Context, tagID int, childIDs []int) error { + if _, err := dbWrapper.Exec(ctx, "DELETE FROM tags_relations WHERE parent_id = ?", tagID); err != nil { + return err + } + + if len(childIDs) > 0 { + var args []interface{} + var values []string + for _, childID := range childIDs { + values = append(values, "(? , ?)") + args = append(args, tagID, childID) + } + + query := "INSERT INTO tags_relations (parent_id, child_id) VALUES " + strings.Join(values, ", ") + if _, err := dbWrapper.Exec(ctx, query, args...); err != nil { + return err + } + } + + return nil +} + +// FindAllAncestors returns a slice of TagPath objects, representing all +// ancestors of the tag with the provided id. +func (qb *TagStore) FindAllAncestors(ctx context.Context, tagID int, excludeIDs []int) ([]*models.TagPath, error) { + inBinding := getInBinding(len(excludeIDs) + 1) + + query := `WITH RECURSIVE +parents AS ( + SELECT t.id AS parent_id, t.id AS child_id, t.name::text as path FROM tags t WHERE t.id = ? + UNION + SELECT tr.parent_id, tr.child_id, t.name || '->' || p.path as path FROM tags_relations tr INNER JOIN parents p ON p.parent_id = tr.child_id JOIN tags t ON t.id = tr.parent_id WHERE tr.parent_id NOT IN` + inBinding + ` +) +SELECT t.*, p.path FROM tags t INNER JOIN parents p ON t.id = p.parent_id +` + + args := []interface{}{tagID, tagID} + for _, excludeID := range excludeIDs { + args = append(args, excludeID) + } + + return qb.queryTagPaths(ctx, query, args) +} + +// FindAllDescendants returns a slice of TagPath objects, representing all +// descendants of the tag with the provided id. +func (qb *TagStore) FindAllDescendants(ctx context.Context, tagID int, excludeIDs []int) ([]*models.TagPath, error) { + inBinding := getInBinding(len(excludeIDs) + 1) + + query := `WITH RECURSIVE +children AS ( + SELECT t.id AS parent_id, t.id AS child_id, t.name::text as path FROM tags t WHERE t.id = ? + UNION + SELECT tr.parent_id, tr.child_id, c.path || '->' || t.name as path FROM tags_relations tr INNER JOIN children c ON c.child_id = tr.parent_id JOIN tags t ON t.id = tr.child_id WHERE tr.child_id NOT IN` + inBinding + ` +) +SELECT t.*, c.path FROM tags t INNER JOIN children c ON t.id = c.child_id +` + + args := []interface{}{tagID, tagID} + for _, excludeID := range excludeIDs { + args = append(args, excludeID) + } + + return qb.queryTagPaths(ctx, query, args) +} + +type tagRelationshipStore struct { + idRelationshipStore +} + +func (s *tagRelationshipStore) CountByTagID(ctx context.Context, tagID int) (int, error) { + joinTable := s.joinTable.table.table + q := dialect.Select(goqu.COUNT("*")).From(joinTable).Where(joinTable.Col(tagIDColumn).Eq(tagID)) + return count(ctx, q) +} + +func (s *tagRelationshipStore) GetTagIDs(ctx context.Context, id int) ([]int, error) { + return s.joinTable.get(ctx, id) +} diff --git a/pkg/postgres/tag_filter.go b/pkg/postgres/tag_filter.go new file mode 100644 index 0000000000..0c4193464b --- /dev/null +++ b/pkg/postgres/tag_filter.go @@ -0,0 +1,236 @@ +package postgres + +import ( + "context" + + "github.com/stashapp/stash/pkg/models" +) + +type tagFilterHandler struct { + tagFilter *models.TagFilterType +} + +func (qb *tagFilterHandler) validate() error { + tagFilter := qb.tagFilter + if tagFilter == nil { + return nil + } + + if err := validateFilterCombination(tagFilter.OperatorFilter); err != nil { + return err + } + + if subFilter := tagFilter.SubFilter(); subFilter != nil { + sqb := &tagFilterHandler{tagFilter: subFilter} + if err := sqb.validate(); err != nil { + return err + } + } + + return nil +} + +func (qb *tagFilterHandler) handle(ctx context.Context, f *filterBuilder) { + tagFilter := qb.tagFilter + if tagFilter == nil { + return + } + + if err := qb.validate(); err != nil { + f.setError(err) + return + } + + sf := tagFilter.SubFilter() + if sf != nil { + sub := &tagFilterHandler{sf} + handleSubFilter(ctx, sub, f, tagFilter.OperatorFilter) + } + + f.handleCriterion(ctx, qb.criterionHandler()) +} + +var tagHierarchyHandler = hierarchicalRelationshipHandler{ + primaryTable: tagTable, + relationTable: tagRelationsTable, + aliasPrefix: tagTable, + parentIDCol: "parent_id", + childIDCol: "child_id", +} + +func (qb *tagFilterHandler) criterionHandler() criterionHandler { + tagFilter := qb.tagFilter + return compoundHandler{ + stringCriterionHandler(tagFilter.Name, tagTable+".name"), + stringCriterionHandler(tagFilter.SortName, tagTable+".sort_name"), + qb.aliasCriterionHandler(tagFilter.Aliases), + + boolCriterionHandler(tagFilter.Favorite, tagTable+".favorite", nil), + stringCriterionHandler(tagFilter.Description, tagTable+".description"), + boolCriterionHandler(tagFilter.IgnoreAutoTag, tagTable+".ignore_auto_tag", nil), + + qb.isMissingCriterionHandler(tagFilter.IsMissing), + qb.sceneCountCriterionHandler(tagFilter.SceneCount), + qb.imageCountCriterionHandler(tagFilter.ImageCount), + qb.galleryCountCriterionHandler(tagFilter.GalleryCount), + qb.performerCountCriterionHandler(tagFilter.PerformerCount), + qb.studioCountCriterionHandler(tagFilter.StudioCount), + + qb.groupCountCriterionHandler(tagFilter.GroupCount), + qb.groupCountCriterionHandler(tagFilter.MovieCount), + + qb.markerCountCriterionHandler(tagFilter.MarkerCount), + tagHierarchyHandler.ParentsCriterionHandler(tagFilter.Parents), + tagHierarchyHandler.ChildrenCriterionHandler(tagFilter.Children), + tagHierarchyHandler.ParentCountCriterionHandler(tagFilter.ParentCount), + tagHierarchyHandler.ChildCountCriterionHandler(tagFilter.ChildCount), + + &stashIDCriterionHandler{ + c: tagFilter.StashIDEndpoint, + stashIDRepository: &tagRepository.stashIDs, + stashIDTableAs: "tag_stash_ids", + parentIDCol: "tags.id", + }, + &stashIDsCriterionHandler{ + c: tagFilter.StashIDsEndpoint, + stashIDRepository: &tagRepository.stashIDs, + stashIDTableAs: "tag_stash_ids", + parentIDCol: "tags.id", + }, + + ×tampCriterionHandler{tagFilter.CreatedAt, "tags.created_at", nil}, + ×tampCriterionHandler{tagFilter.UpdatedAt, "tags.updated_at", nil}, + + &relatedFilterHandler{ + relatedIDCol: "scenes_tags.scene_id", + relatedRepo: sceneRepository.repository, + relatedHandler: &sceneFilterHandler{tagFilter.ScenesFilter}, + joinFn: func(f *filterBuilder) { + tagRepository.scenes.innerJoin(f, "", "tags.id") + }, + }, + + &relatedFilterHandler{ + relatedIDCol: "images_tags.image_id", + relatedRepo: imageRepository.repository, + relatedHandler: &imageFilterHandler{tagFilter.ImagesFilter}, + joinFn: func(f *filterBuilder) { + tagRepository.images.innerJoin(f, "", "tags.id") + }, + }, + + &relatedFilterHandler{ + relatedIDCol: "galleries_tags.gallery_id", + relatedRepo: galleryRepository.repository, + relatedHandler: &galleryFilterHandler{tagFilter.GalleriesFilter}, + joinFn: func(f *filterBuilder) { + tagRepository.galleries.innerJoin(f, "", "tags.id") + }, + }, + } +} + +func (qb *tagFilterHandler) aliasCriterionHandler(alias *models.StringCriterionInput) criterionHandlerFunc { + h := stringListCriterionHandlerBuilder{ + primaryTable: tagTable, + primaryFK: tagIDColumn, + joinTable: tagAliasesTable, + stringColumn: tagAliasColumn, + addJoinTable: func(f *filterBuilder) { + tagRepository.aliases.join(f, "", "tags.id") + }, + } + + return h.handler(alias) +} + +func (qb *tagFilterHandler) isMissingCriterionHandler(isMissing *string) criterionHandlerFunc { + return func(ctx context.Context, f *filterBuilder) { + if isMissing != nil && *isMissing != "" { + switch *isMissing { + case "image": + f.addWhere("tags.image_blob IS NULL") + default: + f.addWhere("(tags." + *isMissing + " IS NULL OR TRIM(CAST(tags." + *isMissing + " AS TEXT)) = '')") + } + } + } +} + +func (qb *tagFilterHandler) sceneCountCriterionHandler(sceneCount *models.IntCriterionInput) criterionHandlerFunc { + return func(ctx context.Context, f *filterBuilder) { + if sceneCount != nil { + f.addLeftJoin("scenes_tags", "", "scenes_tags.tag_id = tags.id") + clause, args := getIntCriterionWhereClause("count(distinct scenes_tags.scene_id)", *sceneCount) + + f.addHaving(clause, args...) + } + } +} + +func (qb *tagFilterHandler) imageCountCriterionHandler(imageCount *models.IntCriterionInput) criterionHandlerFunc { + return func(ctx context.Context, f *filterBuilder) { + if imageCount != nil { + f.addLeftJoin("images_tags", "", "images_tags.tag_id = tags.id") + clause, args := getIntCriterionWhereClause("count(distinct images_tags.image_id)", *imageCount) + + f.addHaving(clause, args...) + } + } +} + +func (qb *tagFilterHandler) galleryCountCriterionHandler(galleryCount *models.IntCriterionInput) criterionHandlerFunc { + return func(ctx context.Context, f *filterBuilder) { + if galleryCount != nil { + f.addLeftJoin("galleries_tags", "", "galleries_tags.tag_id = tags.id") + clause, args := getIntCriterionWhereClause("count(distinct galleries_tags.gallery_id)", *galleryCount) + + f.addHaving(clause, args...) + } + } +} + +func (qb *tagFilterHandler) performerCountCriterionHandler(performerCount *models.IntCriterionInput) criterionHandlerFunc { + return func(ctx context.Context, f *filterBuilder) { + if performerCount != nil { + f.addLeftJoin("performers_tags", "", "performers_tags.tag_id = tags.id") + clause, args := getIntCriterionWhereClause("count(distinct performers_tags.performer_id)", *performerCount) + + f.addHaving(clause, args...) + } + } +} + +func (qb *tagFilterHandler) studioCountCriterionHandler(studioCount *models.IntCriterionInput) criterionHandlerFunc { + return func(ctx context.Context, f *filterBuilder) { + if studioCount != nil { + f.addLeftJoin("studios_tags", "", "studios_tags.tag_id = tags.id") + clause, args := getIntCriterionWhereClause("count(distinct studios_tags.studio_id)", *studioCount) + + f.addHaving(clause, args...) + } + } +} + +func (qb *tagFilterHandler) groupCountCriterionHandler(groupCount *models.IntCriterionInput) criterionHandlerFunc { + return func(ctx context.Context, f *filterBuilder) { + if groupCount != nil { + f.addLeftJoin("groups_tags", "", "groups_tags.tag_id = tags.id") + clause, args := getIntCriterionWhereClause("count(distinct groups_tags.group_id)", *groupCount) + + f.addHaving(clause, args...) + } + } +} + +func (qb *tagFilterHandler) markerCountCriterionHandler(markerCount *models.IntCriterionInput) criterionHandlerFunc { + return func(ctx context.Context, f *filterBuilder) { + if markerCount != nil { + f.addLeftJoin("scene_markers_tags", "", "scene_markers_tags.tag_id = tags.id") + f.addLeftJoin("scene_markers", "", "scene_markers_tags.scene_marker_id = scene_markers.id OR scene_markers.primary_tag_id = tags.id") + clause, args := getIntCriterionWhereClause("count(distinct scene_markers.id)", *markerCount) + + f.addHaving(clause, args...) + } + } +} diff --git a/pkg/postgres/timestamp.go b/pkg/postgres/timestamp.go new file mode 100644 index 0000000000..9b3170fc0d --- /dev/null +++ b/pkg/postgres/timestamp.go @@ -0,0 +1,80 @@ +package postgres + +import ( + "database/sql/driver" + "time" +) + +const TimestampFormat = time.RFC3339 + +// Timestamp represents a time stored in RFC3339 format. +type Timestamp struct { + Timestamp time.Time +} + +// Scan implements the Scanner interface. +func (t *Timestamp) Scan(value interface{}) error { + t.Timestamp = value.(time.Time) + return nil +} + +// Value implements the driver Valuer interface. +func (t Timestamp) Value() (driver.Value, error) { + return t.Timestamp.Format(TimestampFormat), nil +} + +// UTCTimestamp stores a time in UTC. +// TODO - Timestamp should use UTC by default +type UTCTimestamp struct { + Timestamp +} + +// Value implements the driver Valuer interface. +func (t UTCTimestamp) Value() (driver.Value, error) { + return t.Timestamp.Timestamp.UTC().Format(TimestampFormat), nil +} + +// NullTimestamp represents a nullable time stored in RFC3339 format. +type NullTimestamp struct { + Timestamp time.Time + Valid bool +} + +// Scan implements the Scanner interface. +func (t *NullTimestamp) Scan(value interface{}) error { + var ok bool + t.Timestamp, ok = value.(time.Time) + if !ok { + t.Timestamp = time.Time{} + t.Valid = false + return nil + } + + t.Valid = true + return nil +} + +// Value implements the driver Valuer interface. +func (t NullTimestamp) Value() (driver.Value, error) { + if !t.Valid { + return nil, nil + } + + return t.Timestamp.Format(TimestampFormat), nil +} + +func (t NullTimestamp) TimePtr() *time.Time { + if !t.Valid { + return nil + } + + timestamp := t.Timestamp + return ×tamp +} + +func NullTimestampFromTimePtr(t *time.Time) NullTimestamp { + if t == nil { + return NullTimestamp{Valid: false} + } + return NullTimestamp{Timestamp: *t, Valid: true} +} diff --git a/pkg/postgres/transaction.go b/pkg/postgres/transaction.go new file mode 100644 index 0000000000..aebd3cc82e --- /dev/null +++ b/pkg/postgres/transaction.go @@ -0,0 +1,141 @@ +package postgres + +import ( + "context" + "errors" + "fmt" + "runtime/debug" + + "github.com/jackc/pgx/v5/pgconn" + "github.com/jmoiron/sqlx" + "github.com/stashapp/stash/pkg/logger" + "github.com/stashapp/stash/pkg/models" +) + +type key int + +const ( + txnKey key = iota + 1 + dbKey + writableKey +) + +func (db *Database) WithDatabase(ctx context.Context) (context.Context, error) { + // if we are already in a transaction or have a database already, just use it + if tx, _ := getDBReader(ctx); tx != nil { + return ctx, nil + } + + return context.WithValue(ctx, dbKey, db.readDB), nil +} + +func (db *Database) Begin(ctx context.Context, writable bool) (context.Context, error) { + if tx, _ := getTx(ctx); tx != nil { + // log the stack trace so we can see + logger.Error(string(debug.Stack())) + + return nil, fmt.Errorf("already in transaction") + } + + dbtx := db.readDB + if writable { + dbtx = db.writeDB + } + + // Refresh database connection + if err := dbtx.PingContext(ctx); err != nil { + return nil, fmt.Errorf("ping: %w", err) + } + + tx, err := dbtx.BeginTxx(ctx, nil) + if err != nil { + return nil, fmt.Errorf("beginning transaction: %w", err) + } + + ctx = context.WithValue(ctx, writableKey, writable) + + return context.WithValue(ctx, txnKey, tx), nil +} + +func (db *Database) Commit(ctx context.Context) error { + tx, err := getTx(ctx) + if err != nil { + return err + } + + defer db.txnComplete(ctx) + + if err := tx.Commit(); err != nil { + return err + } + + return nil +} + +func (db *Database) Rollback(ctx context.Context) error { + tx, err := getTx(ctx) + if err != nil { + return err + } + + defer db.txnComplete(ctx) + + if err := tx.Rollback(); err != nil { + return err + } + + return nil +} + +func (db *Database) txnComplete(ctx context.Context) { +} + +func getTx(ctx context.Context) (*sqlx.Tx, error) { + tx, ok := ctx.Value(txnKey).(*sqlx.Tx) + if !ok || tx == nil { + return nil, fmt.Errorf("not in transaction") + } + return tx, nil +} + +func getDBReader(ctx context.Context) (dbReader, error) { + // get transaction first if present + tx, ok := ctx.Value(txnKey).(*sqlx.Tx) + if !ok || tx == nil { + // try to get database if present + db, ok := ctx.Value(dbKey).(*sqlx.DB) + if !ok || db == nil { + return nil, fmt.Errorf("not in transaction") + } + return db, nil + } + return tx, nil +} + +func (db *Database) IsLocked(err error) bool { + var pgErr *pgconn.PgError + if errors.As(err, &pgErr) { + // Class 53 — Insufficient Resources + return pgErr.Code[:2] == "53" + } + return false +} + +func (db *Database) Repository() models.Repository { + return models.Repository{ + TxnManager: db, + Blob: db.Blobs(), + File: db.File(), + Folder: db.Folder(), + Gallery: db.Gallery(), + GalleryChapter: db.GalleryChapter(), + Image: db.Image(), + Group: db.Group(), + Performer: db.Performer(), + Scene: db.Scene(), + SceneMarker: db.SceneMarker(), + Studio: db.Studio(), + Tag: db.Tag(), + SavedFilter: db.SavedFilter(), + } +} diff --git a/pkg/postgres/tx.go b/pkg/postgres/tx.go new file mode 100644 index 0000000000..b86452d596 --- /dev/null +++ b/pkg/postgres/tx.go @@ -0,0 +1,160 @@ +package postgres + +import ( + "context" + "database/sql" + "fmt" + "time" + + "github.com/jmoiron/sqlx" + "github.com/stashapp/stash/pkg/logger" +) + +const ( + slowLogTime = time.Millisecond * 200 +) + +type dbReader interface { + Get(dest interface{}, query string, args ...interface{}) error + GetContext(ctx context.Context, dest interface{}, query string, args ...interface{}) error + SelectContext(ctx context.Context, dest interface{}, query string, args ...interface{}) error + QueryxContext(ctx context.Context, query string, args ...interface{}) (*sqlx.Rows, error) +} + +type stmt struct { + *sql.Stmt + query string +} + +func logSQL(start time.Time, query string, args ...interface{}) { + since := time.Since(start) + if since >= slowLogTime { + logger.Debugf("SLOW SQL [%v]: %s, args: %v", since, query, args) + } else { + logger.Tracef("SQL [%v]: %s, args: %v", since, query, args) + } +} + +type dbWrapperType struct{} + +var dbWrapper = dbWrapperType{} + +func sqlError(err error, sql string, args ...interface{}) error { + if err == nil { + return nil + } + + return fmt.Errorf("error executing `%s` [%v]: %w", sql, args, err) +} + +func (db *dbWrapperType) Rebind(query string) string { + return sqlx.Rebind(sqlx.DOLLAR, query) +} + +func (db *dbWrapperType) Get(ctx context.Context, dest interface{}, query string, args ...interface{}) error { + query = db.Rebind(query) + tx, err := getDBReader(ctx) + if err != nil { + return sqlError(err, query, args...) + } + + start := time.Now() + err = tx.GetContext(ctx, dest, query, args...) + logSQL(start, query, args...) + + return sqlError(err, query, args...) +} + +func (db *dbWrapperType) Select(ctx context.Context, dest interface{}, query string, args ...interface{}) error { + query = db.Rebind(query) + tx, err := getDBReader(ctx) + if err != nil { + return sqlError(err, query, args...) + } + + start := time.Now() + err = tx.SelectContext(ctx, dest, query, args...) + logSQL(start, query, args...) + + return sqlError(err, query, args...) +} + +func (db *dbWrapperType) Queryx(ctx context.Context, query string, args ...interface{}) (*sqlx.Rows, error) { + query = db.Rebind(query) + tx, err := getDBReader(ctx) + if err != nil { + return nil, sqlError(err, query, args...) + } + + start := time.Now() + ret, err := tx.QueryxContext(ctx, query, args...) + logSQL(start, query, args...) + + return ret, sqlError(err, query, args...) +} + +func (db *dbWrapperType) QueryxContext(ctx context.Context, query string, args ...interface{}) (*sqlx.Rows, error) { + query = db.Rebind(query) + return dbWrapper.Queryx(ctx, query, args...) +} + +func (db *dbWrapperType) NamedExec(ctx context.Context, query string, arg interface{}) (sql.Result, error) { + query = db.Rebind(query) + tx, err := getTx(ctx) + if err != nil { + return nil, sqlError(err, query, arg) + } + + start := time.Now() + ret, err := tx.NamedExecContext(ctx, query, arg) + logSQL(start, query, arg) + + return ret, sqlError(err, query, arg) +} + +func (db *dbWrapperType) Exec(ctx context.Context, query string, args ...interface{}) (sql.Result, error) { + query = db.Rebind(query) + tx, err := getTx(ctx) + if err != nil { + return nil, sqlError(err, query, args...) + } + + start := time.Now() + ret, err := tx.ExecContext(ctx, query, args...) + logSQL(start, query, args...) + + return ret, sqlError(err, query, args...) +} + +// Prepare creates a prepared statement. +func (db *dbWrapperType) Prepare(ctx context.Context, query string, args ...interface{}) (*stmt, error) { + query = db.Rebind(query) + tx, err := getTx(ctx) + if err != nil { + return nil, sqlError(err, query, args...) + } + + // nolint:sqlclosecheck + ret, err := tx.PrepareContext(ctx, query) + if err != nil { + return nil, sqlError(err, query, args...) + } + + return &stmt{ + query: query, + Stmt: ret, + }, nil +} + +func (*dbWrapperType) ExecStmt(ctx context.Context, stmt *stmt, args ...interface{}) (sql.Result, error) { + _, err := getTx(ctx) + if err != nil { + return nil, sqlError(err, stmt.query, args...) + } + + start := time.Now() + ret, err := stmt.ExecContext(ctx, args...) + logSQL(start, stmt.query, args...) + + return ret, sqlError(err, stmt.query, args...) +} diff --git a/pkg/postgres/values.go b/pkg/postgres/values.go new file mode 100644 index 0000000000..e14fc08f93 --- /dev/null +++ b/pkg/postgres/values.go @@ -0,0 +1,70 @@ +package postgres + +import ( + "gopkg.in/guregu/null.v4" + + "github.com/stashapp/stash/pkg/models" +) + +// null package does not provide methods to convert null.Int to int pointer +func intFromPtr(i *int) null.Int { + if i == nil { + return null.NewInt(0, false) + } + + return null.IntFrom(int64(*i)) +} + +func nullIntPtr(i null.Int) *int { + if !i.Valid { + return nil + } + + v := int(i.Int64) + return &v +} + +func nullFloatPtr(i null.Float) *float64 { + if !i.Valid { + return nil + } + + v := float64(i.Float64) + return &v +} + +func nullIntFolderIDPtr(i null.Int) *models.FolderID { + if !i.Valid { + return nil + } + + v := models.FolderID(i.Int64) + + return &v +} + +func nullIntFileIDPtr(i null.Int) *models.FileID { + if !i.Valid { + return nil + } + + v := models.FileID(i.Int64) + + return &v +} + +func nullIntFromFileIDPtr(i *models.FileID) null.Int { + if i == nil { + return null.NewInt(0, false) + } + + return null.IntFrom(int64(*i)) +} + +func nullIntFromFolderIDPtr(i *models.FolderID) null.Int { + if i == nil { + return null.NewInt(0, false) + } + + return null.IntFrom(int64(*i)) +} diff --git a/pkg/scene/export_test.go b/pkg/scene/export_test.go index cde421bd80..9ca84f6726 100644 --- a/pkg/scene/export_test.go +++ b/pkg/scene/export_test.go @@ -68,7 +68,7 @@ var names = []string{ var imageBytes = []byte("imageBytes") var stashID = models.StashID{ - StashID: "StashID", + StashID: getUUID("StashID"), Endpoint: "Endpoint", } diff --git a/pkg/scene/import_test.go b/pkg/scene/import_test.go index a6e3edcdfd..f4d54cc9e4 100644 --- a/pkg/scene/import_test.go +++ b/pkg/scene/import_test.go @@ -42,6 +42,11 @@ var ( var testCtx = context.Background() +func getUUID(_ string) string { + // TODO: Encode input string + return "00000000-0000-0000-0000-000000000000" +} + func TestImporterPreImport(t *testing.T) { var ( title = "title" @@ -49,7 +54,7 @@ func TestImporterPreImport(t *testing.T) { details = "details" director = "director" endpoint1 = "endpoint1" - stashID1 = "stashID1" + stashID1 = getUUID("stashID1") endpoint2 = "endpoint2" stashID2 = "stashID2" url1 = "url1" diff --git a/pkg/scene/update_test.go b/pkg/scene/update_test.go index f72c964039..6e83292a32 100644 --- a/pkg/scene/update_test.go +++ b/pkg/scene/update_test.go @@ -107,7 +107,7 @@ func TestUpdater_Update(t *testing.T) { performerIDs := []int{performerID} tagIDs := []int{tagID} - stashID := "stashID" + stashID := getUUID("stashID") endpoint := "endpoint" title := "title" @@ -235,7 +235,7 @@ func TestUpdateSet_UpdateInput(t *testing.T) { performerIDStrs := intslice.IntSliceToStringSlice(performerIDs) tagIDs := []int{tagID} tagIDStrs := intslice.IntSliceToStringSlice(tagIDs) - stashID := "stashID" + stashID := getUUID("stashID") endpoint := "endpoint" updatedAt := time.Now() stashIDs := []models.StashID{ diff --git a/pkg/sqlite/anonymise.go b/pkg/sqlite/anonymise.go index 764f569c01..7b4f3e717f 100644 --- a/pkg/sqlite/anonymise.go +++ b/pkg/sqlite/anonymise.go @@ -40,6 +40,28 @@ func NewAnonymiser(db *Database, outPath string) (*Anonymiser, error) { return &Anonymiser{Database: newDB}, nil } +type ForeignAnonymiser interface { + FetchAll(ctx context.Context) error + GetSqliteDatabase() *Database +} + +func PassAnonymiser(sourceDB ForeignAnonymiser) (*Anonymiser, error) { + db := sourceDB.GetSqliteDatabase() + + db.writeDB.Close() + + db.writeDB, _ = db.open(true, true) + db.writeDB.SetMaxOpenConns(1) + db.writeDB.SetMaxIdleConns(10) + db.writeDB.SetConnMaxIdleTime(dbConnTimeout) + + if err := sourceDB.FetchAll(context.Background()); err != nil { + return nil, fmt.Errorf("fetching postgres: %w", err) + } + + return &Anonymiser{Database: db}, nil +} + func (db *Anonymiser) Anonymise(ctx context.Context) error { if err := func() error { defer db.Close() diff --git a/pkg/sqlite/batch.go b/pkg/sqlite/batch.go index a594388356..6494b9ebb2 100644 --- a/pkg/sqlite/batch.go +++ b/pkg/sqlite/batch.go @@ -1,20 +1,11 @@ package sqlite -const defaultBatchSize = 1000 +import ( + "github.com/stashapp/stash/pkg/database" +) -// batchExec executes the provided function in batches of the provided size. -func batchExec[T any](ids []T, batchSize int, fn func(batch []T) error) error { - for i := 0; i < len(ids); i += batchSize { - end := i + batchSize - if end > len(ids) { - end = len(ids) - } - - batch := ids[i:end] - if err := fn(batch); err != nil { - return err - } - } +const defaultBatchSize = database.DefaultBatchSize - return nil +func batchExec[T any](ids []T, batchSize int, fn func(batch []T) error) error { + return database.BatchExec(ids, batchSize, fn) } diff --git a/pkg/sqlite/blob.go b/pkg/sqlite/blob.go index 241b63d23c..0a92aa2e3f 100644 --- a/pkg/sqlite/blob.go +++ b/pkg/sqlite/blob.go @@ -11,6 +11,7 @@ import ( "github.com/doug-martin/goqu/v9/exp" "github.com/jmoiron/sqlx" "github.com/mattn/go-sqlite3" + "github.com/stashapp/stash/pkg/database" "github.com/stashapp/stash/pkg/file" "github.com/stashapp/stash/pkg/hash/md5" "github.com/stashapp/stash/pkg/logger" @@ -24,18 +25,6 @@ const ( blobChecksumColumn = "checksum" ) -type BlobStoreOptions struct { - // UseFilesystem should be true if blob data should be stored in the filesystem - UseFilesystem bool - // UseDatabase should be true if blob data should be stored in the database - UseDatabase bool - // Path is the filesystem path to use for storing blobs - Path string - // SupplementaryPaths are alternative filesystem paths that will be used to find blobs - // No changes will be made to these filesystems - SupplementaryPaths []string -} - type BlobStore struct { repository @@ -44,10 +33,10 @@ type BlobStore struct { fsStore *blob.FilesystemStore // supplementary stores otherStores []blob.FilesystemReader - options BlobStoreOptions + options database.BlobStoreOptions } -func NewBlobStore(options BlobStoreOptions) *BlobStore { +func NewBlobStore(options database.BlobStoreOptions) *BlobStore { fs := &file.OsFS{} ret := &BlobStore{ diff --git a/pkg/sqlite/custom_fields.go b/pkg/sqlite/custom_fields.go index 63f85b250f..b5bce95486 100644 --- a/pkg/sqlite/custom_fields.go +++ b/pkg/sqlite/custom_fields.go @@ -3,6 +3,7 @@ package sqlite import ( "context" "fmt" + "reflect" "regexp" "strings" @@ -130,7 +131,7 @@ func (s *customFieldsStore) setCustomFields(ctx context.Context, id int, values conflictKey := s.fk.GetCol().(string) + ", field" // upsert new custom fields - q := dialect.Insert(s.table).Prepared(true).Cols(s.fk, "field", "value"). + q := dialect.Insert(s.table).Prepared(true).Cols(s.fk, "field", "value", "type"). OnConflict(goqu.DoUpdate(conflictKey, goqu.Record{"value": goqu.I("excluded.value")})) r := make([]interface{}, len(values)) var i int @@ -139,7 +140,7 @@ func (s *customFieldsStore) setCustomFields(ctx context.Context, id int, values if err != nil { return fmt.Errorf("getting SQL value for field %q: %w", key, err) } - r[i] = goqu.Record{"field": key, "value": v, s.fk.GetCol().(string): id} + r[i] = goqu.Record{"field": key, "value": v, "type": reflect.TypeOf(v).String(), s.fk.GetCol().(string): id} i++ } diff --git a/pkg/sqlite/database.go b/pkg/sqlite/database.go index 0ea3d71700..1e23da5686 100644 --- a/pkg/sqlite/database.go +++ b/pkg/sqlite/database.go @@ -8,10 +8,12 @@ import ( "fmt" "os" "path/filepath" + "strings" "time" "github.com/jmoiron/sqlx" + "github.com/stashapp/stash/pkg/database" "github.com/stashapp/stash/pkg/fsutil" "github.com/stashapp/stash/pkg/logger" ) @@ -34,37 +36,11 @@ const ( cacheSizeEnv = "STASH_SQLITE_CACHE_SIZE" ) -var appSchemaVersion uint = 75 +var appSchemaVersion uint = 76 //go:embed migrations/*.sql var migrationsBox embed.FS -var ( - // ErrDatabaseNotInitialized indicates that the database is not - // initialized, usually due to an incomplete configuration. - ErrDatabaseNotInitialized = errors.New("database not initialized") -) - -// ErrMigrationNeeded indicates that a database migration is needed -// before the database can be initialized -type MigrationNeededError struct { - CurrentSchemaVersion uint - RequiredSchemaVersion uint -} - -func (e *MigrationNeededError) Error() string { - return fmt.Sprintf("database schema version %d does not match required schema version %d", e.CurrentSchemaVersion, e.RequiredSchemaVersion) -} - -type MismatchedSchemaVersionError struct { - CurrentSchemaVersion uint - RequiredSchemaVersion uint -} - -func (e *MismatchedSchemaVersionError) Error() string { - return fmt.Sprintf("schema version %d is incompatible with required schema version %d", e.CurrentSchemaVersion, e.RequiredSchemaVersion) -} - type storeRepository struct { Blobs *BlobStore File *FileStore @@ -97,7 +73,7 @@ func NewDatabase() *Database { fileStore := NewFileStore() folderStore := NewFolderStore() galleryStore := NewGalleryStore(fileStore, folderStore) - blobStore := NewBlobStore(BlobStoreOptions{}) + blobStore := NewBlobStore(database.BlobStoreOptions{}) performerStore := NewPerformerStore(blobStore) studioStore := NewStudioStore(blobStore) tagStore := NewTagStore(blobStore) @@ -127,14 +103,18 @@ func NewDatabase() *Database { return ret } -func (db *Database) SetBlobStoreOptions(options BlobStoreOptions) { - *db.Blobs = *NewBlobStore(options) +func (db *Database) DatabaseBackend() database.DatabaseType { + return database.SqliteBackend +} + +func (db *Database) SetBlobStoreOptions(options database.BlobStoreOptions) { + *db.storeRepository.Blobs = *NewBlobStore(options) } // Ready returns an error if the database is not ready to begin transactions. func (db *Database) Ready() error { if db.readDB == nil || db.writeDB == nil { - return ErrDatabaseNotInitialized + return database.ErrDatabaseNotInitialized } return nil @@ -148,7 +128,7 @@ func (db *Database) Open(dbPath string) error { db.lock() defer db.unlock() - db.dbPath = dbPath + db.dbPath, _ = strings.CutPrefix(dbPath, string(database.SqliteBackend)+":") databaseSchemaVersion, err := db.getDatabaseSchemaVersion() if err != nil { @@ -166,7 +146,7 @@ func (db *Database) Open(dbPath string) error { } } else { if databaseSchemaVersion > appSchemaVersion { - return &MismatchedSchemaVersionError{ + return &database.MismatchedSchemaVersionError{ CurrentSchemaVersion: databaseSchemaVersion, RequiredSchemaVersion: appSchemaVersion, } @@ -174,7 +154,7 @@ func (db *Database) Open(dbPath string) error { // if migration is needed, then don't open the connection if db.needsMigration() { - return &MigrationNeededError{ + return &database.MigrationNeededError{ CurrentSchemaVersion: databaseSchemaVersion, RequiredSchemaVersion: appSchemaVersion, } diff --git a/pkg/sqlite/date.go b/pkg/sqlite/date.go index 522fe7cb09..41c4a15936 100644 --- a/pkg/sqlite/date.go +++ b/pkg/sqlite/date.go @@ -1,79 +1,19 @@ package sqlite import ( - "database/sql/driver" - "time" - + "github.com/stashapp/stash/pkg/database" "github.com/stashapp/stash/pkg/models" "gopkg.in/guregu/null.v4" ) -const sqliteDateLayout = "2006-01-02" - // Date represents a date stored as "YYYY-MM-DD" -type Date struct { - Date time.Time -} - -// Scan implements the Scanner interface. -func (d *Date) Scan(value interface{}) error { - d.Date = value.(time.Time) - return nil -} - -// Value implements the driver Valuer interface. -func (d Date) Value() (driver.Value, error) { - return d.Date.Format(sqliteDateLayout), nil -} +type Date = database.Date // NullDate represents a nullable date stored as "YYYY-MM-DD" -type NullDate struct { - Date time.Time - Valid bool -} - -// Scan implements the Scanner interface. -func (d *NullDate) Scan(value interface{}) error { - var ok bool - d.Date, ok = value.(time.Time) - if !ok { - d.Date = time.Time{} - d.Valid = false - return nil - } +type NullDate = database.NullDate - d.Valid = true - return nil -} - -// Value implements the driver Valuer interface. -func (d NullDate) Value() (driver.Value, error) { - if !d.Valid { - return nil, nil - } - - return d.Date.Format(sqliteDateLayout), nil -} - -func (d *NullDate) DatePtr(precision null.Int) *models.Date { - if d == nil || !d.Valid { - return nil - } - - return &models.Date{Time: d.Date, Precision: models.DatePrecision(precision.Int64)} -} - -func NullDateFromDatePtr(d *models.Date) NullDate { - if d == nil { - return NullDate{Valid: false} - } - return NullDate{Date: d.Time, Valid: true} -} +var NullDateFromDatePtr = database.NullDateFromDatePtr func datePrecisionFromDatePtr(d *models.Date) null.Int { - if d == nil { - // default to day precision - return null.Int{} - } - return null.IntFrom(int64(d.Precision)) + return database.DatePrecisionFromDatePtr(d) } diff --git a/pkg/sqlite/interfaces.go b/pkg/sqlite/interfaces.go new file mode 100644 index 0000000000..9a2e166a62 --- /dev/null +++ b/pkg/sqlite/interfaces.go @@ -0,0 +1,46 @@ +package sqlite + +import "github.com/stashapp/stash/pkg/database" + +func (db *Database) Blobs() database.BlobStore { + return db.storeRepository.Blobs +} +func (db *Database) File() database.FileStore { + return db.storeRepository.File +} +func (db *Database) Folder() database.FolderStore { + return db.storeRepository.Folder +} +func (db *Database) Image() database.ImageStore { + return db.storeRepository.Image +} +func (db *Database) Gallery() database.GalleryStore { + return db.storeRepository.Gallery +} +func (db *Database) GalleryChapter() database.GalleryChapterStore { + return db.storeRepository.GalleryChapter +} +func (db *Database) Scene() database.SceneStore { + return db.storeRepository.Scene +} +func (db *Database) SceneMarker() database.SceneMarkerStore { + return db.storeRepository.SceneMarker +} +func (db *Database) Performer() database.PerformerStore { + return db.storeRepository.Performer +} +func (db *Database) SavedFilter() database.SavedFilterStore { + return db.storeRepository.SavedFilter +} +func (db *Database) Studio() database.StudioStore { + return db.storeRepository.Studio +} +func (db *Database) Tag() database.TagStore { + return db.storeRepository.Tag +} +func (db *Database) Group() database.GroupStore { + return db.storeRepository.Group +} +func (db *Database) NewMigrator() (database.MigrateStore, error) { + return NewMigrator(db) +} diff --git a/pkg/sqlite/migrations/76_cf_type_field.up.sql b/pkg/sqlite/migrations/76_cf_type_field.up.sql new file mode 100644 index 0000000000..ac2d946989 --- /dev/null +++ b/pkg/sqlite/migrations/76_cf_type_field.up.sql @@ -0,0 +1,2 @@ +-- 73_cf_type_field.up.sql +ALTER TABLE performer_custom_fields ADD COLUMN `type` text; diff --git a/pkg/sqlite/migrations/76_postmigrate.go b/pkg/sqlite/migrations/76_postmigrate.go new file mode 100644 index 0000000000..c51b364c1c --- /dev/null +++ b/pkg/sqlite/migrations/76_postmigrate.go @@ -0,0 +1,77 @@ +package migrations + +import ( + "context" + "fmt" + "reflect" + + "github.com/jmoiron/sqlx" + "github.com/stashapp/stash/pkg/logger" + "github.com/stashapp/stash/pkg/sqlite" +) + +type schema73Migrator struct { + migrator +} + +func post73(ctx context.Context, db *sqlx.DB) error { + logger.Info("Running post-migration for schema version 73") + + m := schema73Migrator{ + migrator: migrator{ + db: db, + }, + } + + return m.migrate(ctx) +} + +func (m *schema73Migrator) migrate(ctx context.Context) error { + if err := m.withTxn(ctx, func(tx *sqlx.Tx) error { + query := "SELECT `performer_id`, `value` FROM `performer_custom_fields`" + + rows, err := tx.Queryx(query) + if err != nil { + return err + } + defer rows.Close() + + for rows.Next() { + var ( + performer_id int + value interface{} + ) + + err := rows.Scan(&performer_id, &value) + if err != nil { + return err + } + + gotype := reflect.TypeOf(value).String() + logger.Debugf("setting type for %v to %v for %v", value, gotype, performer_id) + r, err := tx.Exec("UPDATE performer_custom_fields SET type = ? WHERE performer_id = ?", gotype, performer_id) + if err != nil { + return fmt.Errorf("error setting type for %v to %v for %v", value, gotype, performer_id) + } + + rowsAffected, err := r.RowsAffected() + if err != nil { + return err + } + + if rowsAffected == 0 { + return fmt.Errorf("no rows affected when updating type %v to %v for %v", value, gotype, performer_id) + } + } + + return rows.Err() + }); err != nil { + return err + } + + return nil +} + +func init() { + sqlite.RegisterPostMigration(73, post73) +} diff --git a/pkg/sqlite/sql.go b/pkg/sqlite/sql.go index 0b55af8db4..b3f15c2341 100644 --- a/pkg/sqlite/sql.go +++ b/pkg/sqlite/sql.go @@ -130,7 +130,7 @@ func getRandomSort(tableName string, direction string, seed uint64) string { // ORDER BY ((n+seed)*(n+seed)*p1 + (n+seed)*p2) % p3 // since sqlite converts overflowing numbers to reals, a custom db function that uses uints with overflow should be faster, // however in practice the overhead of calling a custom function vastly outweighs the benefits - return fmt.Sprintf(" ORDER BY mod((%[1]s + %[2]d) * (%[1]s + %[2]d) * 52959209 + (%[1]s + %[2]d) * 1047483763, 2147483647) %[3]s", colName, seed, direction) + return fmt.Sprintf(" ORDER BY (%[1]s + %[2]d) * (%[1]s + %[2]d) * 52959209 + (%[1]s + %[2]d) * 1047483763 %% 2147483647 %[3]s", colName, seed, direction) } func getCountSort(primaryTable, joinTable, primaryFK, direction string) string { diff --git a/pkg/sqlite/transaction.go b/pkg/sqlite/transaction.go index fb86723bdf..9ca8883d4f 100644 --- a/pkg/sqlite/transaction.go +++ b/pkg/sqlite/transaction.go @@ -118,18 +118,18 @@ func (db *Database) IsLocked(err error) bool { func (db *Database) Repository() models.Repository { return models.Repository{ TxnManager: db, - Blob: db.Blobs, - File: db.File, - Folder: db.Folder, - Gallery: db.Gallery, - GalleryChapter: db.GalleryChapter, - Image: db.Image, - Group: db.Group, - Performer: db.Performer, - Scene: db.Scene, - SceneMarker: db.SceneMarker, - Studio: db.Studio, - Tag: db.Tag, - SavedFilter: db.SavedFilter, + Blob: db.Blobs(), + File: db.File(), + Folder: db.Folder(), + Gallery: db.Gallery(), + GalleryChapter: db.GalleryChapter(), + Image: db.Image(), + Group: db.Group(), + Performer: db.Performer(), + Scene: db.Scene(), + SceneMarker: db.SceneMarker(), + Studio: db.Studio(), + Tag: db.Tag(), + SavedFilter: db.SavedFilter(), } } diff --git a/pkg/studio/export_test.go b/pkg/studio/export_test.go index c333c0ad5b..4e9ed6e966 100644 --- a/pkg/studio/export_test.go +++ b/pkg/studio/export_test.go @@ -40,9 +40,14 @@ var parentStudio models.Studio = models.Studio{ var imageBytes = []byte("imageBytes") +func getUUID(_ string) string { + // TODO: Encode input string + return "00000000-0000-0000-0000-000000000000" +} + var aliases = []string{"alias"} var stashID = models.StashID{ - StashID: "StashID", + StashID: getUUID("StashID"), Endpoint: "Endpoint", } var stashIDs = []models.StashID{