diff --git a/.gitignore b/.gitignore index 93f2dc8e..98521fd9 100644 --- a/.gitignore +++ b/.gitignore @@ -16,3 +16,4 @@ custom/ *.secret gokapi-cli gokapi-cli.json +test_debian13.txt diff --git a/build/go.mod b/build/go.mod index 2d6860db..a04080e9 100644 --- a/build/go.mod +++ b/build/go.mod @@ -8,7 +8,9 @@ require ( github.com/alicebob/miniredis/v2 v2.38.0 github.com/aws/aws-sdk-go v1.55.8 github.com/caarlos0/env/v6 v6.10.1 + github.com/go-sql-driver/mysql v1.10.1 github.com/gomodule/redigo v1.9.3 + github.com/jackc/pgx/v5 v5.10.0 github.com/jinzhu/copier v0.4.0 github.com/johannesboyne/gofakes3 v1.2.0 github.com/juju/ratelimit v1.0.2 @@ -26,12 +28,17 @@ require ( ) require ( + filippo.io/edwards25519 v1.2.0 // indirect github.com/dustin/go-humanize v1.0.1 // indirect github.com/ebitengine/purego v0.10.1 // indirect github.com/go-jose/go-jose/v4 v4.1.4 // indirect github.com/go-ole/go-ole v1.3.0 // indirect github.com/google/uuid v1.6.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/jmespath/go-jmespath v0.4.0 // indirect + github.com/kr/text v0.1.0 // indirect github.com/lufia/plan9stats v0.0.0-20260627054121-477a66015f15 // indirect github.com/mattn/go-isatty v0.0.22 // indirect github.com/mitchellh/colorstring v0.0.0-20190213212951-d06e56a500db // indirect @@ -39,6 +46,7 @@ require ( github.com/power-devops/perfstat v0.0.0-20240221224432-82ca36839d55 // indirect github.com/remyoudompheng/bigfft v0.0.0-20230129092748-24d4a6f8daec // indirect github.com/rivo/uniseg v0.4.7 // indirect + github.com/rogpeppe/go-internal v1.16.0 // indirect github.com/ryszard/goskiplist v0.0.0-20150312221310-2dfbae5fcf46 // indirect github.com/tdewolff/minify/v2 v2.24.11 // indirect github.com/tdewolff/parse/v2 v2.8.11 // indirect @@ -48,6 +56,7 @@ require ( github.com/yusufpapurcu/wmi v1.2.4 // indirect go.shabbyrobe.org/gocovmerge v0.0.0-20230507111327-fa4f82cfbf4d // indirect golang.org/x/exp v0.0.0-20251023183803-a4bb9ffd2546 // indirect + golang.org/x/text v0.40.0 // indirect golang.org/x/tools v0.48.0 // indirect gopkg.in/check.v1 v1.0.0-20201130134442-10cb98267c6c // indirect modernc.org/libc v1.74.1 // indirect diff --git a/build/go.sum b/build/go.sum index 5aa07796..c2174bf2 100644 --- a/build/go.sum +++ b/build/go.sum @@ -1,32 +1,32 @@ +filippo.io/edwards25519 v1.2.0/go.mod h1:xzAOLCNug/yB62zG1bQ8uziwrIqIuxhctzJT18Q77mc= github.com/Kodeworks/golang-image-ico v0.0.0-20141118225523-73f0f4cfade9/go.mod h1:7uhhqiBaR4CpN0k9rMjOtjpcfGd6DG2m04zQxKnWQ0I= github.com/NYTimes/gziphandler v1.1.1/go.mod h1:n/CVRwUEOgIxrgPvAQhUUr9oeUtvrhMomdKFjzJNB0c= -github.com/alicebob/miniredis/v2 v2.37.0/go.mod h1:TcL7YfarKPGDAthEtl5NBeHZfeUQj6OXMm/+iu5cLMM= github.com/alicebob/miniredis/v2 v2.38.0/go.mod h1:TcL7YfarKPGDAthEtl5NBeHZfeUQj6OXMm/+iu5cLMM= github.com/aws/aws-sdk-go v1.55.8/go.mod h1:ZkViS9AqA6otK+JBBNH2++sx1sgxrPKcSzPPvQkUtXk= github.com/caarlos0/env/v6 v6.10.1/go.mod h1:hvp/ryKXKipEkcuYjs9mI4bBCg+UI0Yhgm5Zu0ddvwc= -github.com/coreos/go-oidc/v3 v3.17.0/go.mod h1:wqPbKFrVnE90vty060SB40FCJ8fTHTxSwyXJqZH+sI8= github.com/coreos/go-oidc/v3 v3.20.0/go.mod h1:DYCf24+ncYi+XkIH97GY1+dqoRlbaSI26KVTCI9SrY4= github.com/davecgh/go-spew v1.1.0/go.mod h1:J7Y8YcW2NihsgmVo/mv3lAwl/skON4iLHjSsI+c5H38= github.com/dustin/go-humanize v1.0.1/go.mod h1:Mu1zIs6XwVuF/gI1OepvI0qD18qycQx+mFykh5fBlto= -github.com/ebitengine/purego v0.10.0/go.mod h1:iIjxzd6CiRiOG0UyXP+V1+jWqUXVjPKLAI0mRfJZTmQ= github.com/ebitengine/purego v0.10.1/go.mod h1:iIjxzd6CiRiOG0UyXP+V1+jWqUXVjPKLAI0mRfJZTmQ= github.com/go-jose/go-jose/v4 v4.1.4/go.mod h1:x4oUasVrzR7071A4TnHLGSPpNOm2a21K9Kf04k1rs08= github.com/go-ole/go-ole v1.2.6/go.mod h1:pprOEPIfldk/42T2oK7lQ4v4JSDwmV0As9GaiUsvbm0= github.com/go-ole/go-ole v1.3.0/go.mod h1:5LS6F96DhAwUc7C+1HLexzMXY1xGRSryjyPPKW6zv78= +github.com/go-sql-driver/mysql v1.10.1/go.mod h1:M+cqaI7+xxXGG9swrdeUIoPG3Y3KCkF0pZej+SK+nWk= github.com/gomodule/redigo v1.9.3/go.mod h1:KsU3hiK/Ay8U42qpaJk+kuNa3C+spxapWpM+ywhcgtw= github.com/google/uuid v1.6.0/go.mod h1:TIyPZe4MgqvfeYDBFedMoGGpEw/LqOeaOT+nhxU+yHo= +github.com/jackc/pgpassfile v1.0.0/go.mod h1:CEx0iS5ambNFdcRtxPj5JhEz+xB6uRky5eyVu/W2HEg= +github.com/jackc/pgservicefile v0.0.0-20240606120523-5a60cdf6a761/go.mod h1:5TJZWKEWniPve33vlWYSoGYefn3gLQRzjfDlhSJ9ZKM= +github.com/jackc/pgx/v5 v5.10.0/go.mod h1:mal1tBGAFfLHvZzaYh77YS/eC6IX9OWbRV1QIIM0Jn4= +github.com/jackc/puddle/v2 v2.2.2/go.mod h1:vriiEXHvEE654aYKXXjOvZM39qJ0q+azkZFrfEOc3H4= github.com/jinzhu/copier v0.4.0/go.mod h1:DfbEm0FYsaqBcKcFuvmOZb218JkPGtvSHsKg8S8hyyg= github.com/jmespath/go-jmespath v0.4.0/go.mod h1:T8mJZnbsbmF+m6zOOFylbeCJqk5+pHWvzYPziyZiYoo= github.com/jmespath/go-jmespath/internal/testify v1.5.1/go.mod h1:L3OGu8Wl2/fWfCI6z80xFu9LTZmf1ZRjMHUOPmWr69U= -github.com/johannesboyne/gofakes3 v0.0.0-20260208201424-4c385a1f6a73/go.mod h1:S4S9jGBVlLri0OeqrSSbCGG5vsI6he06UJyuz1WT1EE= github.com/johannesboyne/gofakes3 v1.2.0/go.mod h1:UHhRZRod9rENGFrUWTYnQHZqlNgSmjOq8DaD/ATQYRM= github.com/juju/ratelimit v1.0.2/go.mod h1:qapgC/Gy+xNh9UxzV13HGGl/6UXNN+ct+vwSgWNm/qk= github.com/kr/pretty v0.2.1/go.mod h1:ipq/a2n7PKx3OHsz4KJII5eveXtPO4qwEXGdVfWzfnI= github.com/kr/pty v1.1.1/go.mod h1:pFQYn66WHrOpPYNljwOMqo10TkYh1fy3cYio2l3bCsQ= github.com/kr/text v0.1.0/go.mod h1:4Jbv+DJW3UT/LiOwJeYQe1efqtUx/iVham/4vfdArNI= -github.com/lufia/plan9stats v0.0.0-20260330125221-c963978e514e/go.mod h1:autxFIvghDt3jPTLoqZ9OZ7s9qTGNAWmYCjVFWPX/zg= github.com/lufia/plan9stats v0.0.0-20260627054121-477a66015f15/go.mod h1:autxFIvghDt3jPTLoqZ9OZ7s9qTGNAWmYCjVFWPX/zg= -github.com/mattn/go-isatty v0.0.20/go.mod h1:W+V8PltTTMOvKvAeJH7IuucS94S2C6jfK/D7dTCTo3Y= github.com/mattn/go-isatty v0.0.22/go.mod h1:ZXfXG4SQHsB/w3ZeOYbR0PrPwLy+n6xiMrJlRFqopa4= github.com/mitchellh/colorstring v0.0.0-20190213212951-d06e56a500db/go.mod h1:l0dey0ia/Uv7NcFFVbCLtqEBQbrT4OCwCSKTEv6enCw= github.com/ncruces/go-strftime v1.0.0/go.mod h1:Fwc5htZGVVkseilnfgOVb9mKy6w1naJmn9CehxcKcls= @@ -34,34 +34,40 @@ github.com/pmezard/go-difflib v1.0.0/go.mod h1:iKH77koFhYxTK1pcRnkKkqfTogsbg7gZN github.com/power-devops/perfstat v0.0.0-20240221224432-82ca36839d55/go.mod h1:OmDBASR4679mdNQnz2pUhc2G8CO2JrUAVFDRBDP/hJE= github.com/remyoudompheng/bigfft v0.0.0-20230129092748-24d4a6f8daec/go.mod h1:qqbHyh8v60DhA7CoWK5oRCqLrMHRGoxYCSS9EjAz6Eo= github.com/rivo/uniseg v0.4.7/go.mod h1:FN3SvrM+Zdj16jyLfmOkMNblXMcoc8DfTHruCPUcx88= +github.com/rogpeppe/go-internal v1.16.0/go.mod h1:DrUVZyrJU+txYW5/1kwtXQSMFio52ZOxX7yM1VHvnxs= github.com/ryszard/goskiplist v0.0.0-20150312221310-2dfbae5fcf46/go.mod h1:uAQ5PCi+MFsC7HjREoAz1BU+Mq60+05gifQSsHSDG/8= -github.com/schollz/progressbar/v3 v3.19.0/go.mod h1:IsO3lpbaGuzh8zIMzgY3+J8l4C8GjO0Y9S69eFvNsec= github.com/schollz/progressbar/v3 v3.19.1/go.mod h1:LFL7jqimKxfhero4K1eCkUr/6R39AgQeiPCJtlTWIW8= github.com/secure-io/sio-go v0.3.1/go.mod h1:+xbkjDzPjwh4Axd07pRKSNriS9SCiYksWnZqdnfpQxs= -github.com/shirou/gopsutil/v4 v4.26.3/go.mod h1:LZ6ewCSkBqUpvSOf+LsTGnRinC6iaNUNMGBtDkJBaLQ= github.com/shirou/gopsutil/v4 v4.26.6/go.mod h1:LZ6ewCSkBqUpvSOf+LsTGnRinC6iaNUNMGBtDkJBaLQ= github.com/stretchr/objx v0.1.0/go.mod h1:HFkY916IF+rwdDfMAkV7OtwuqBVzrE8GR6GFx+wExME= github.com/stretchr/testify v1.3.0/go.mod h1:M5WIy9Dh21IEIfnGCwXGc5bZfKNJtfHm1UVUgZn+9EI= +github.com/stretchr/testify v1.7.0/go.mod h1:6Fq8oRcR53rry900zMqJjRRixrwX3KX962/h/Wwjteg= +github.com/tdewolff/minify/v2 v2.24.11 h1:JlANsiWaRBXedoYtsiZgY3YFkdr42oF32vp2SLgQKi4= github.com/tdewolff/minify/v2 v2.24.11/go.mod h1:exq1pjdrh9uAICdfVKQwqz6MsJmWmQahZuTC6pTO6ro= +github.com/tdewolff/minify/v2 v2.24.17 h1:6AbitfVyq0M7aW6i+XL7+49DeTQZwloOMs9O574arBg= +github.com/tdewolff/minify/v2 v2.24.17/go.mod h1:kVqn9vxXUKtlHexSNrWbYePqioOT5mc4ou/KVSMpfCM= +github.com/tdewolff/parse/v2 v2.8.11 h1:SGyjEy3xEqd+W9WVzTlTQ5GkP/en4a1AZNZVJ1cvgm0= github.com/tdewolff/parse/v2 v2.8.11/go.mod h1:Hwlni2tiVNKyzR1o6nUs4FOF07URA+JLBLd6dlIXYqo= +github.com/tdewolff/parse/v2 v2.8.16 h1:bLk5svUOQRkW/Y2SJ+DeENSIkZBcTIkq+Atyv5D8feI= +github.com/tdewolff/parse/v2 v2.8.16/go.mod h1:XdsoSFThlVIRIajAuqz1evNY7bagZS8LBOPA3aVopwQ= github.com/tdewolff/test v1.0.11/go.mod h1:XPuWBzvdUzhCuxWO1ojpXsyzsA5bFoS3tO/Q3kFuTG8= -github.com/tklauser/go-sysconf v0.3.16/go.mod h1:/qNL9xxDhc7tx3HSRsLWNnuzbVfh3e7gh/BmM179nYI= +github.com/tdewolff/test v1.0.12 h1:7F21DqIajswxuche0geHdrUZRCWE4oko4b7bcmkkrxk= +github.com/tdewolff/test v1.0.12/go.mod h1:XPuWBzvdUzhCuxWO1ojpXsyzsA5bFoS3tO/Q3kFuTG8= github.com/tklauser/go-sysconf v0.4.0/go.mod h1:8mTNWyog7H+MpKijp4VmKJAd2bbYQ2zuUwkYRbUArPI= -github.com/tklauser/numcpus v0.11.0/go.mod h1:z+LwcLq54uWZTX0u/bGobaV34u6V7KNlTZejzM6/3MQ= github.com/tklauser/numcpus v0.12.0/go.mod h1:ABHeXzJnr/qqwguhClkZKT1/8VABcYrsyUiUGobwWJg= github.com/yuin/gopher-lua v1.1.2/go.mod h1:7aRmXIWl37SqRf0koeyylBEzJ+aPt8A+mmkQ4f1ntR8= github.com/yusufpapurcu/wmi v1.2.4/go.mod h1:SBZ9tNy3G9/m5Oi98Zks0QjeHVDvuK0qfxQmPyzfmi0= go.shabbyrobe.org/gocovmerge v0.0.0-20230507111327-fa4f82cfbf4d/go.mod h1:92Uoe3l++MlthCm+koNi0tcUCX3anayogF0Pa/sp24k= golang.org/x/crypto v0.0.0-20190308221718-c2843e01d9a2/go.mod h1:djNgcEr1/C05ACkg1iLfiJU5Ep61QUkGW8qpdssI0+w= golang.org/x/crypto v0.0.0-20200302210943-78000ba7a073/go.mod h1:LzIPMQfyMNhhGPhUkYOs5KpL4U8rLKemX1yGLhDgUto= -golang.org/x/crypto v0.49.0/go.mod h1:ErX4dUh2UM+CFYiXZRTcMpEcN8b/1gxEuv3nODoYtCA= golang.org/x/crypto v0.54.0/go.mod h1:KWL8ny2AZdGR2cWmzeHrp2azQPGogOv+HeQaVEXC2dk= +golang.org/x/exp v0.0.0-20251023183803-a4bb9ffd2546 h1:mgKeJMpvi0yx/sU5GsxQ7p6s2wtOnGAHZWCHUM4KGzY= golang.org/x/exp v0.0.0-20251023183803-a4bb9ffd2546/go.mod h1:j/pmGrbnkbPtQfxEe5D0VQhZC6qKbfKifgD0oM7sR70= -golang.org/x/image v0.38.0/go.mod h1:/3f6vaXC+6CEanU4KJxbcUZyEePbyKbaLoDOe4ehFYY= +golang.org/x/exp v0.0.0-20260908205506-85c1c2202aba h1:Ck8QetSgk912qxWLMCKxd0in+aiyBQyDSMae6e/xmpU= +golang.org/x/exp v0.0.0-20260908205506-85c1c2202aba/go.mod h1:50RgIsmK7OwqzTTeqcSXQW8SswW0o8fRcDxmqGluJ8E= golang.org/x/image v0.44.0/go.mod h1:V8K3KE9KKKE+pLpQDOeN18w9oacNSvy1tDOirTu4xtY= golang.org/x/net v0.0.0-20190404232315-eb5bcb51f2a3/go.mod h1:t9HGtf8HONx5eT2rtn7q6eTqICYqUVnKs3thJo3Qplg= golang.org/x/oauth2 v0.36.0/go.mod h1:YDBUJMTkDnJS+A4BP4eZBjCqtokkg1hODuPjwiGPO7Q= -golang.org/x/sync v0.20.0/go.mod h1:9xrNwdLfx4jkKbNva9FpL6vEN7evnE43NNNJQ2LF3+0= golang.org/x/sync v0.22.0/go.mod h1:9xrNwdLfx4jkKbNva9FpL6vEN7evnE43NNNJQ2LF3+0= golang.org/x/sys v0.0.0-20190215142949-d0b11bdaac8a/go.mod h1:STP8DvDyc/dI5b8T5hshtkjS+E42TnysNCUPdjciGhY= golang.org/x/sys v0.0.0-20190412213103-97732733099d/go.mod h1:h1NjWce9XRLGQEsW7wpKNCjG9DtNlClVuFLEZdDNbEs= @@ -69,22 +75,19 @@ golang.org/x/sys v0.0.0-20190916202348-b4ddaad3f8a3/go.mod h1:h1NjWce9XRLGQEsW7w golang.org/x/sys v0.0.0-20200302150141-5c8b2ff67527/go.mod h1:h1NjWce9XRLGQEsW7wpKNCjG9DtNlClVuFLEZdDNbEs= golang.org/x/sys v0.0.0-20201204225414-ed752295db88/go.mod h1:h1NjWce9XRLGQEsW7wpKNCjG9DtNlClVuFLEZdDNbEs= golang.org/x/sys v0.1.0/go.mod h1:oPkhp1MJrh7nUepCBck5+mAzfO9JrbApNNgaTdGDITg= -golang.org/x/sys v0.6.0/go.mod h1:oPkhp1MJrh7nUepCBck5+mAzfO9JrbApNNgaTdGDITg= -golang.org/x/sys v0.42.0/go.mod h1:4GL1E5IUh+htKOUEOaiffhrAeqysfVGipDYzABqnCmw= golang.org/x/sys v0.47.0/go.mod h1:4GL1E5IUh+htKOUEOaiffhrAeqysfVGipDYzABqnCmw= -golang.org/x/term v0.41.0/go.mod h1:3pfBgksrReYfZ5lvYM0kSO0LIkAl4Yl2bXOkKP7Ec2A= golang.org/x/term v0.45.0/go.mod h1:9aqxs0blBcrm/n0L9QW0aRVD+ktan8ssZromtqJC43w= golang.org/x/text v0.3.0/go.mod h1:NqM8EUOU14njkJ3fqMW+pc6Ldnwhi/IjpwHt7yyuwOQ= +golang.org/x/text v0.40.0/go.mod h1:hpnzDAfGV753zIKo+wk3u1bVKCGPbrnF7+7LBF/UHVY= golang.org/x/time v0.15.0/go.mod h1:Y4YMaQmXwGQZoFaVFk4YpCt4FLQMYKZe9oeV/f4MSno= -golang.org/x/tools v0.43.0/go.mod h1:uHkMso649BX2cZK6+RpuIPXS3ho2hZo4FVwfoy1vIk0= golang.org/x/tools v0.48.0/go.mod h1:08xX0orndb/F7jJxGDicx061tyd5pcMto75YMAXr6lk= gopkg.in/check.v1 v0.0.0-20161208181325-20d25e280405/go.mod h1:Co6ibVJAznAaIkqp8huTwlJQCZ016jof/cbN4VW5Yz0= +gopkg.in/check.v1 v1.0.0-20201130134442-10cb98267c6c h1:Hei/4ADfdWqJk1ZMxUNpqntNwaWcugrBjAiHlqqRiVk= gopkg.in/check.v1 v1.0.0-20201130134442-10cb98267c6c/go.mod h1:JHkPIbrfpd72SG/EVd6muEfDQjcINNoR0C8j2r3qZ4Q= gopkg.in/yaml.v2 v2.2.8/go.mod h1:hI93XBmqTisBFMUTm0b8Fm+jr3Dg1NNxqwp+5A1VGuI= +gopkg.in/yaml.v3 v3.0.0-20200313102051-9f266ea9e77c/go.mod h1:K4uyk7z7BCEPqu6E+C64Yfv1cQ7kz7rIZviUmN+EgEM= gopkg.in/yaml.v3 v3.0.1/go.mod h1:K4uyk7z7BCEPqu6E+C64Yfv1cQ7kz7rIZviUmN+EgEM= -modernc.org/libc v1.70.0/go.mod h1:OVmxFGP1CI/Z4L3E0Q3Mf1PDE0BucwMkcXjjLntvHJo= modernc.org/libc v1.74.1/go.mod h1:uH4t5bOx3G3g9Xcmj10YKlTcVISlRDwv8VoQJG9n8Os= modernc.org/mathutil v1.7.1/go.mod h1:4p5IwJITfppl0G4sUEDtCr4DthTaT47/N3aT6MhfgJg= modernc.org/memory v1.11.0/go.mod h1:/JP4VbVC+K5sU2wZi9bHoq2MAkCnrt2r98UGeSK7Mjw= -modernc.org/sqlite v1.48.1/go.mod h1:hWjRO6Tj/5Ik8ieqxQybiEOUXy0NJFNp2tpvVpKlvig= modernc.org/sqlite v1.53.0/go.mod h1:xoEpOIpGrgT48H5iiyt/YXPCZPEzlfmfFwtk8Lklw8s= diff --git a/docs/advanced.rst b/docs/advanced.rst index 952d4eeb..44ce4991 100644 --- a/docs/advanced.rst +++ b/docs/advanced.rst @@ -216,13 +216,15 @@ For Docker users, the command is: Database URL format --------------------------------- -Database URLs must start with either ``sqlite://`` or ``redis://``. +Database URLs must start with ``sqlite://``, ``redis://``, ``mariadb://`` or ``postgres://``. For SQLite, the path to the database follows the prefix. No additional options are allowed. For Redis, the URL can include authentication credentials (username and password), an optional prefix for keys, and parameter to use SSL. +For MariaDB/MySQL and PostgreSQL, the URL can include authentication credentials and must include the database name. The target database must already exist. + Redis URL Format --------------------------------- @@ -231,7 +233,7 @@ A Redis URL has the following structure: :: redis://[username:password@]host[:port][?options] - + * username: (optional) The username for authentication. * password: (optional) The password for authentication. * host: (required) The address of the Redis server. @@ -239,6 +241,36 @@ A Redis URL has the following structure: * options: (optional) Additional options such as SSL (``ssl=true``) and key prefix (``prefix=``). +MariaDB / MySQL URL Format +--------------------------------- + +A MariaDB/MySQL URL has the following structure: +:: + + mariadb://[username:password@]host[:port]/database + +* username: (optional) The username for authentication. +* password: (optional) The password for authentication. +* host: (required) The address of the MariaDB/MySQL server. +* port: (optional) The port of the server (default is 3306). +* database: (required) The name of the database to use. + + +PostgreSQL URL Format +--------------------------------- + +A PostgreSQL URL has the following structure: +:: + + postgres://[username:password@]host[:port]/database + +* username: (optional) The username for authentication. +* password: (optional) The password for authentication. +* host: (required) The address of the PostgreSQL server. +* port: (optional) The port of the server (default is 5432). +* database: (required) The name of the database to use. + + Examples --------------------------------- @@ -262,6 +294,20 @@ Migrating Redis (``127.0.0.1:6379, User: test, Password: 1234, Prefix: gokapi_, gokapi migrate-database --source "redis://test:1234@127.0.0.1:6379?prefix=gokapi_&ssl=true" --destination sqlite://./data/gokapi.sqlite +Migrating SQLite (``./data/gokapi.sqlite``) to PostgreSQL (``127.0.0.1:5432, Database: gokapi, User: gokapi, Password: secret``): + + +:: + + gokapi migrate-database --source sqlite://./data/gokapi.sqlite --destination "postgres://gokapi:secret@127.0.0.1:5432/gokapi" + +Migrating MariaDB (``127.0.0.1:3306, Database: gokapi, User: gokapi, Password: secret``) to PostgreSQL (``127.0.0.1:5432, Database: gokapi, User: gokapi, Password: secret``): + + +:: + + gokapi migrate-database --source "mariadb://gokapi:secret@127.0.0.1:3306/gokapi" --destination "postgres://gokapi:secret@127.0.0.1:5432/gokapi" + .. _clitool: diff --git a/docs/setup.rst b/docs/setup.rst index 3da070a8..025f2950 100644 --- a/docs/setup.rst +++ b/docs/setup.rst @@ -133,20 +133,21 @@ Database If you choose Redis, **you must enable Redis persistence** before storing any data (e.g. add ``save 1 1`` to your ``redis.conf``). Without persistence, all data is lost on a Redis restart. .. warning:: - The Redis password is stored in plain text in the configuration file and will be visible if you re-run setup. + The Redis, MariaDB/MySQL, and PostgreSQL passwords are stored in plain text in the configuration file and will be visible if you re-run setup. -By default Gokapi uses SQLite, which is fine for most deployments. Use Redis if: +By default Gokapi uses SQLite, which is fine for most deployments. Use Redis, MariaDB/MySQL, or PostgreSQL if: * you expect high download/upload traffic, or * your SQLite database lives on a slow disk (e.g. a network share or SD card). Settings: -* **Type of database** — SQLite or Redis. +* **Type of database** — SQLite, Redis, MariaDB/MySQL, or PostgreSQL. * **Database location** — path to the SQLite file. -* **Database host** — host and port for Redis (e.g. ``127.0.0.1:6379``). +* **Database host** — host and port for Redis (e.g. ``127.0.0.1:6379``), MariaDB/MySQL (e.g. ``127.0.0.1:3306``), or PostgreSQL (e.g. ``127.0.0.1:5432``). * **Key prefix** *(optional)* — added to all Redis keys; useful when sharing a Redis instance with other applications. -* **Username / Password** *(optional)* — Redis authentication credentials. +* **Database name** — the MariaDB/MySQL or PostgreSQL database to use; it must already exist. +* **Username / Password** — Redis authentication is optional; MariaDB/MySQL and PostgreSQL require a username (password is technically optional but strongly recommended). * **Use SSL** — enables TLS for the Redis connection. .. _setup_webserver: diff --git a/go.mod b/go.mod index 54ca6e3f..f44202eb 100644 --- a/go.mod +++ b/go.mod @@ -8,7 +8,9 @@ require ( github.com/alicebob/miniredis/v2 v2.39.0 github.com/aws/aws-sdk-go v1.55.8 github.com/caarlos0/env/v6 v6.10.1 + github.com/go-sql-driver/mysql v1.10.1 github.com/gomodule/redigo v1.9.3 + github.com/jackc/pgx/v5 v5.10.0 github.com/jinzhu/copier v0.4.0 github.com/johannesboyne/gofakes3 v1.2.0 github.com/juju/ratelimit v1.0.2 @@ -26,12 +28,17 @@ require ( ) require ( + filippo.io/edwards25519 v1.2.0 // indirect github.com/dustin/go-humanize v1.0.1 // indirect github.com/ebitengine/purego v0.10.2 // indirect github.com/go-jose/go-jose/v4 v4.1.4 // indirect github.com/go-ole/go-ole v1.3.0 // indirect github.com/google/uuid v1.6.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/jmespath/go-jmespath v0.4.0 // indirect + github.com/kr/text v0.1.0 // indirect github.com/lufia/plan9stats v0.0.0-20260627054121-477a66015f15 // indirect github.com/mattn/go-isatty v0.0.24 // indirect github.com/mitchellh/colorstring v0.0.0-20190213212951-d06e56a500db // indirect @@ -39,6 +46,7 @@ require ( github.com/power-devops/perfstat v0.0.0-20240221224432-82ca36839d55 // indirect github.com/remyoudompheng/bigfft v0.0.0-20230129092748-24d4a6f8daec // indirect github.com/rivo/uniseg v0.4.7 // indirect + github.com/rogpeppe/go-internal v1.16.0 // indirect github.com/ryszard/goskiplist v0.0.0-20150312221310-2dfbae5fcf46 // indirect github.com/tdewolff/minify/v2 v2.24.17 // indirect github.com/tdewolff/parse/v2 v2.8.16 // indirect @@ -48,8 +56,8 @@ require ( github.com/yusufpapurcu/wmi v1.2.4 // indirect go.shabbyrobe.org/gocovmerge v0.0.0-20230507111327-fa4f82cfbf4d // indirect golang.org/x/exp v0.0.0-20260908205506-85c1c2202aba // indirect + golang.org/x/text v0.42.0 // indirect golang.org/x/tools v0.50.0 // indirect - gopkg.in/check.v1 v1.0.0-20201130134442-10cb98267c6c // indirect modernc.org/libc v1.75.7 // indirect modernc.org/mathutil v1.7.1 // indirect modernc.org/memory v1.12.1 // indirect diff --git a/go.sum b/go.sum index 2094043d..c123a79e 100644 --- a/go.sum +++ b/go.sum @@ -1,3 +1,5 @@ +filippo.io/edwards25519 v1.2.0 h1:crnVqOiS4jqYleHd9vaKZ+HKtHfllngJIiOpNpoJsjo= +filippo.io/edwards25519 v1.2.0/go.mod h1:xzAOLCNug/yB62zG1bQ8uziwrIqIuxhctzJT18Q77mc= github.com/Kodeworks/golang-image-ico v0.0.0-20141118225523-73f0f4cfade9 h1:1ltqoej5GtaWF8jaiA49HwsZD459jqm9YFz9ZtMFpQA= github.com/Kodeworks/golang-image-ico v0.0.0-20141118225523-73f0f4cfade9/go.mod h1:7uhhqiBaR4CpN0k9rMjOtjpcfGd6DG2m04zQxKnWQ0I= github.com/NYTimes/gziphandler v1.1.1 h1:ZUDjpQae29j0ryrS0u/B8HZfJBtBQHjqw2rQ2cqUQ3I= @@ -52,6 +54,8 @@ github.com/go-jose/go-jose/v4 v4.1.4/go.mod h1:x4oUasVrzR7071A4TnHLGSPpNOm2a21K9 github.com/go-ole/go-ole v1.2.6/go.mod h1:pprOEPIfldk/42T2oK7lQ4v4JSDwmV0As9GaiUsvbm0= github.com/go-ole/go-ole v1.3.0 h1:Dt6ye7+vXGIKZ7Xtk4s6/xVdGDQynvom7xCFEdWr6uE= github.com/go-ole/go-ole v1.3.0/go.mod h1:5LS6F96DhAwUc7C+1HLexzMXY1xGRSryjyPPKW6zv78= +github.com/go-sql-driver/mysql v1.10.1 h1:arlSnNLq6a5yxGxV7qg9lF4j0C+KwD6NbQyKr9QL6ME= +github.com/go-sql-driver/mysql v1.10.1/go.mod h1:M+cqaI7+xxXGG9swrdeUIoPG3Y3KCkF0pZej+SK+nWk= github.com/gomodule/redigo v1.9.3 h1:dNPSXeXv6HCq2jdyWfjgmhBdqnR6PRO3m/G05nvpPC8= github.com/gomodule/redigo v1.9.3/go.mod h1:KsU3hiK/Ay8U42qpaJk+kuNa3C+spxapWpM+ywhcgtw= github.com/google/go-cmp v0.7.0 h1:wk8382ETsv4JYUZwIsn6YpYiWiBsYLSJiTsyBybVuN8= @@ -62,6 +66,14 @@ github.com/google/uuid v1.6.0 h1:NIvaJDMOsjHA8n1jAhLSgzrAzy1Hgr+hNrb57e+94F0= github.com/google/uuid v1.6.0/go.mod h1:TIyPZe4MgqvfeYDBFedMoGGpEw/LqOeaOT+nhxU+yHo= github.com/hashicorp/golang-lru/v2 v2.0.7 h1:a+bsQ5rvGLjzHuww6tVxozPZFVghXaHOwFs4luLUK2k= github.com/hashicorp/golang-lru/v2 v2.0.7/go.mod h1:QeFd9opnmA6QUJc5vARoKUSoFhyfM2/ZepoAG6RGpeM= +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.10.0 h1:VhSvgU2jSli8o3AqIEOTJr7rZwAEUVo4E4XhR94Zfr0= +github.com/jackc/pgx/v5 v5.10.0/go.mod h1:mal1tBGAFfLHvZzaYh77YS/eC6IX9OWbRV1QIIM0Jn4= +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/jmespath/go-jmespath v0.4.0 h1:BEgLn5cpjn8UN1mAw4NjwDrS35OdebyEtFe+9YPoQUg= @@ -72,8 +84,8 @@ github.com/johannesboyne/gofakes3 v1.2.0 h1:I9VEzPWvvAUAGzDlhYFoZjF0AXMlkcEyZlmB github.com/johannesboyne/gofakes3 v1.2.0/go.mod h1:UHhRZRod9rENGFrUWTYnQHZqlNgSmjOq8DaD/ATQYRM= github.com/juju/ratelimit v1.0.2 h1:sRxmtRiajbvrcLQT7S+JbqU0ntsb9W2yhSdNN8tWfaI= github.com/juju/ratelimit v1.0.2/go.mod h1:qapgC/Gy+xNh9UxzV13HGGl/6UXNN+ct+vwSgWNm/qk= -github.com/kr/pretty v0.2.1 h1:Fmg33tUaq4/8ym9TJN1x7sLJnHVwhP33CNkpYV/7rwI= -github.com/kr/pretty v0.2.1/go.mod h1:ipq/a2n7PKx3OHsz4KJII5eveXtPO4qwEXGdVfWzfnI= +github.com/kr/pretty v0.3.0 h1:WgNl7dwNpEZ6jJ9k1snq4pZsg7DOEN8hP9Xw0Tsjwk0= +github.com/kr/pretty v0.3.0/go.mod h1:640gp4NfQd8pI5XOwp5fnNeVWj67G7CFk/SaSQn7NBk= github.com/kr/pty v1.1.1/go.mod h1:pFQYn66WHrOpPYNljwOMqo10TkYh1fy3cYio2l3bCsQ= github.com/kr/text v0.1.0 h1:45sCR5RtlFHMR4UwH9sdQ5TC8v0qDQCHnXt+kaKSTVE= github.com/kr/text v0.1.0/go.mod h1:4Jbv+DJW3UT/LiOwJeYQe1efqtUx/iVham/4vfdArNI= @@ -95,6 +107,8 @@ github.com/remyoudompheng/bigfft v0.0.0-20230129092748-24d4a6f8daec h1:W09IVJc94 github.com/remyoudompheng/bigfft v0.0.0-20230129092748-24d4a6f8daec/go.mod h1:qqbHyh8v60DhA7CoWK5oRCqLrMHRGoxYCSS9EjAz6Eo= github.com/rivo/uniseg v0.4.7 h1:WUdvkW8uEhrYfLC4ZzdpI2ztxP1I582+49Oc5Mq64VQ= github.com/rivo/uniseg v0.4.7/go.mod h1:FN3SvrM+Zdj16jyLfmOkMNblXMcoc8DfTHruCPUcx88= +github.com/rogpeppe/go-internal v1.16.0 h1:O9DK+vNMDVGLr2BeZqmpLeMjiMNkuXfcqntWbZV6S5g= +github.com/rogpeppe/go-internal v1.16.0/go.mod h1:DrUVZyrJU+txYW5/1kwtXQSMFio52ZOxX7yM1VHvnxs= github.com/ryszard/goskiplist v0.0.0-20150312221310-2dfbae5fcf46 h1:GHRpF1pTW19a8tTFrMLUcfWwyC0pnifVo2ClaLq+hP8= github.com/ryszard/goskiplist v0.0.0-20150312221310-2dfbae5fcf46/go.mod h1:uAQ5PCi+MFsC7HjREoAz1BU+Mq60+05gifQSsHSDG/8= github.com/schollz/progressbar/v3 v3.19.1 h1:iv8BgwOvdML/S3p84uBpy/IMigv4U9594vPZYa2EdrU= @@ -107,6 +121,7 @@ github.com/spf13/afero v1.2.1 h1:qgMbHoJbPbw579P+1zVY+6n4nIFuIchaIjzZ/I/Yq8M= github.com/spf13/afero v1.2.1/go.mod h1:9ZxEEn6pIJ8Rxe320qSDBk6AsU0r9pR7Q4OcevTdifk= github.com/stretchr/objx v0.1.0/go.mod h1:HFkY916IF+rwdDfMAkV7OtwuqBVzrE8GR6GFx+wExME= github.com/stretchr/testify v1.3.0/go.mod h1:M5WIy9Dh21IEIfnGCwXGc5bZfKNJtfHm1UVUgZn+9EI= +github.com/stretchr/testify v1.7.0/go.mod h1:6Fq8oRcR53rry900zMqJjRRixrwX3KX962/h/Wwjteg= github.com/stretchr/testify v1.11.1 h1:7s2iGBzp5EwR7/aIZr8ao5+dra3wiQyKjjFuvgVKu7U= github.com/stretchr/testify v1.11.1/go.mod h1:wZwfW3scLgRK+23gO65QZefKpKQRnfz6sD981Nm4B6U= github.com/tdewolff/minify/v2 v2.24.17 h1:6AbitfVyq0M7aW6i+XL7+49DeTQZwloOMs9O574arBg= @@ -165,6 +180,7 @@ gopkg.in/mgo.v2 v2.0.0-20180705113604-9856a29383ce h1:xcEWjVhvbDy+nHP67nPDDpbYrY gopkg.in/mgo.v2 v2.0.0-20180705113604-9856a29383ce/go.mod h1:yeKp02qBN3iKW1OzL3MGk2IdtZzaj7SFntXj72NppTA= gopkg.in/yaml.v2 v2.2.8 h1:obN1ZagJSUGI0Ek/LBmuj4SNLPfIny3KsKFopxRdj10= gopkg.in/yaml.v2 v2.2.8/go.mod h1:hI93XBmqTisBFMUTm0b8Fm+jr3Dg1NNxqwp+5A1VGuI= +gopkg.in/yaml.v3 v3.0.0-20200313102051-9f266ea9e77c/go.mod h1:K4uyk7z7BCEPqu6E+C64Yfv1cQ7kz7rIZviUmN+EgEM= gopkg.in/yaml.v3 v3.0.1 h1:fxVm/GzAzEWqLHuvctI91KS9hhNmmWOoWu0XTYJS7CA= gopkg.in/yaml.v3 v3.0.1/go.mod h1:K4uyk7z7BCEPqu6E+C64Yfv1cQ7kz7rIZviUmN+EgEM= modernc.org/cc/v4 v4.29.2 h1:h6+9ciCnPKutf4I03CvheAvDLX7+IHlqR6Iy6J+cgd8= diff --git a/internal/configuration/database/Database.go b/internal/configuration/database/Database.go index 6fca4e77..466043d6 100644 --- a/internal/configuration/database/Database.go +++ b/internal/configuration/database/Database.go @@ -51,6 +51,14 @@ func ParseUrl(dbUrl string, mustExist bool) (models.DbConnection, error) { case "redis": result.Type = dbabstraction.TypeRedis result.HostUrl = u.Host + case "mariadb": + result.Type = dbabstraction.TypeMariaDb + result.HostUrl = u.Host + result.DatabaseName = strings.TrimPrefix(u.Path, "/") + case "postgres": + result.Type = dbabstraction.TypePostgres + result.HostUrl = u.Host + result.DatabaseName = strings.TrimPrefix(u.Path, "/") default: return models.DbConnection{}, fmt.Errorf("unsupported database type: %s\n", dbUrl) } diff --git a/internal/configuration/database/Database_test.go b/internal/configuration/database/Database_test.go index 29e6fd28..6c7fae76 100644 --- a/internal/configuration/database/Database_test.go +++ b/internal/configuration/database/Database_test.go @@ -450,6 +450,28 @@ func TestParseUrl(t *testing.T) { output, err = ParseUrl("redis://tuser:tpw@127.0.0.1:1234/?ssl=true&prefix=tpref", false) test.IsNil(t, err) test.IsEqual(t, output, expectedOutput) + + expectedOutput = models.DbConnection{ + HostUrl: "127.0.0.1:3306", + Username: "tuser", + Password: "tpw", + DatabaseName: "gokapi", + Type: dbabstraction.TypeMariaDb, + } + output, err = ParseUrl("mariadb://tuser:tpw@127.0.0.1:3306/gokapi", false) + test.IsNil(t, err) + test.IsEqual(t, output, expectedOutput) + + expectedOutput = models.DbConnection{ + HostUrl: "127.0.0.1:5432", + Username: "tuser", + Password: "tpw", + DatabaseName: "gokapi", + Type: dbabstraction.TypePostgres, + } + output, err = ParseUrl("postgres://tuser:tpw@127.0.0.1:5432/gokapi", false) + test.IsNil(t, err) + test.IsEqual(t, output, expectedOutput) } func TestMigration(t *testing.T) { diff --git a/internal/configuration/database/dbabstraction/DbAbstraction.go b/internal/configuration/database/dbabstraction/DbAbstraction.go index fea4ec20..04f56153 100644 --- a/internal/configuration/database/dbabstraction/DbAbstraction.go +++ b/internal/configuration/database/dbabstraction/DbAbstraction.go @@ -3,6 +3,8 @@ package dbabstraction import ( "fmt" + "github.com/forceu/gokapi/internal/configuration/database/provider/mariadb" + "github.com/forceu/gokapi/internal/configuration/database/provider/postgres" "github.com/forceu/gokapi/internal/configuration/database/provider/redis" "github.com/forceu/gokapi/internal/configuration/database/provider/sqlite" "github.com/forceu/gokapi/internal/models" @@ -13,6 +15,10 @@ const ( TypeSqlite = iota // TypeRedis specifies to use a Redis database TypeRedis + // TypeMariaDb specifies to use a MariaDB or MySQL database + TypeMariaDb + // TypePostgres specifies to use a PostgreSQL database + TypePostgres ) // Database declares the required functions for a database connection @@ -126,6 +132,10 @@ func GetNew(config models.DbConnection) (Database, error) { return sqlite.New(config) case TypeRedis: return redis.New(config) + case TypeMariaDb: + return mariadb.New(config) + case TypePostgres: + return postgres.New(config) default: return nil, fmt.Errorf("unsupported database: type %v", config.Type) } diff --git a/internal/configuration/database/dbabstraction/DbAbstraction_test.go b/internal/configuration/database/dbabstraction/DbAbstraction_test.go index b7ac20c3..4270bf81 100644 --- a/internal/configuration/database/dbabstraction/DbAbstraction_test.go +++ b/internal/configuration/database/dbabstraction/DbAbstraction_test.go @@ -14,6 +14,14 @@ var configRedis = models.DbConnection{ Type: 1, // dbabstraction.TypeRedis } +var configMariaDb = models.DbConnection{ + Type: 2, // dbabstraction.TypeMariaDb +} + +var configPostgres = models.DbConnection{ + Type: 3, // dbabstraction.TypePostgres +} + func TestGetNew(t *testing.T) { result, err := GetNew(configSqlite) test.IsNotNil(t, err) @@ -21,7 +29,13 @@ func TestGetNew(t *testing.T) { result, err = GetNew(configRedis) test.IsNotNil(t, err) test.IsEqualInt(t, result.GetType(), 1) + result, err = GetNew(configMariaDb) + test.IsNotNil(t, err) + test.IsEqualInt(t, result.GetType(), 2) + result, err = GetNew(configPostgres) + test.IsNotNil(t, err) + test.IsEqualInt(t, result.GetType(), 3) - _, err = GetNew(models.DbConnection{Type: 2}) + _, err = GetNew(models.DbConnection{Type: 4}) test.IsNotNil(t, err) } diff --git a/internal/configuration/database/migration/Migration.go b/internal/configuration/database/migration/Migration.go index 463680d2..83b4e52a 100644 --- a/internal/configuration/database/migration/Migration.go +++ b/internal/configuration/database/migration/Migration.go @@ -32,6 +32,10 @@ func getType(input int) string { return "SQLite" case dbabstraction.TypeRedis: return "Redis" + case dbabstraction.TypeMariaDb: + return "MariaDB" + case dbabstraction.TypePostgres: + return "PostgreSQL" } return "Invalid" } diff --git a/internal/configuration/database/migration/Migration_test.go b/internal/configuration/database/migration/Migration_test.go index 0442e5bb..cf36336f 100644 --- a/internal/configuration/database/migration/Migration_test.go +++ b/internal/configuration/database/migration/Migration_test.go @@ -21,7 +21,9 @@ func TestMain(m *testing.M) { func TestGetType(t *testing.T) { test.IsEqualString(t, getType(dbabstraction.TypeSqlite), "SQLite") test.IsEqualString(t, getType(dbabstraction.TypeRedis), "Redis") - test.IsEqualString(t, getType(2), "Invalid") + test.IsEqualString(t, getType(dbabstraction.TypeMariaDb), "MariaDB") + test.IsEqualString(t, getType(dbabstraction.TypePostgres), "PostgreSQL") + test.IsEqualString(t, getType(4), "Invalid") } var exitCode int diff --git a/internal/configuration/database/provider/mariadb/Mariadb.go b/internal/configuration/database/provider/mariadb/Mariadb.go new file mode 100644 index 00000000..14316b7d --- /dev/null +++ b/internal/configuration/database/provider/mariadb/Mariadb.go @@ -0,0 +1,226 @@ +package mariadb + +import ( + "database/sql" + "errors" + "fmt" + + "github.com/forceu/gokapi/internal/helper" + "github.com/forceu/gokapi/internal/models" + // Required for the mariadb/mysql driver + _ "github.com/go-sql-driver/mysql" +) + +// DatabaseProvider contains the database instance +type DatabaseProvider struct { + sqlDb *sql.DB +} + +// DatabaseSchemeVersion contains the version number to be expected from the current database. If lower, an upgrade will be performed +const DatabaseSchemeVersion = 1 + +// New returns an instance +func New(dbConfig models.DbConnection) (DatabaseProvider, error) { + return DatabaseProvider{}.init(dbConfig) +} + +// GetType returns 2, for being a MariaDB/MySQL interface +func (p DatabaseProvider) GetType() int { + return 2 // dbabstraction.TypeMariaDb +} + +// Upgrade migrates the DB to a new Gokapi version, if required +func (p DatabaseProvider) Upgrade(currentDbVersion int) { + // No upgrade steps yet - this provider was introduced at schema version 1 +} + +// GetDbVersion gets the version number of the database +func (p DatabaseProvider) GetDbVersion() int { + var version int + row := p.sqlDb.QueryRow("SELECT Version FROM SchemaVersion LIMIT 1") + err := row.Scan(&version) + if err != nil { + if errors.Is(err, sql.ErrNoRows) { + return 0 + } + helper.Check(err) + } + return version +} + +// SetDbVersion sets the version number of the database +func (p DatabaseProvider) SetDbVersion(newVersion int) { + _, err := p.sqlDb.Exec("DELETE FROM SchemaVersion") + helper.Check(err) + _, err = p.sqlDb.Exec("INSERT INTO SchemaVersion (Version) VALUES (?)", newVersion) + helper.Check(err) +} + +// GetSchemaVersion returns the version number, which the database should be at if fully upgraded +func (p DatabaseProvider) GetSchemaVersion() int { + return DatabaseSchemeVersion +} + +// Init connects to the database and creates the table structure, if necessary +func (p DatabaseProvider) init(dbConfig models.DbConnection) (DatabaseProvider, error) { + if dbConfig.HostUrl == "" { + return DatabaseProvider{}, errors.New("empty database url was provided") + } + if dbConfig.DatabaseName == "" { + return DatabaseProvider{}, errors.New("no database name was provided") + } + if p.sqlDb == nil { + dsn := fmt.Sprintf("%s:%s@tcp(%s)/%s?parseTime=false", dbConfig.Username, dbConfig.Password, dbConfig.HostUrl, dbConfig.DatabaseName) + var err error + p.sqlDb, err = sql.Open("mysql", dsn) + if err != nil { + return DatabaseProvider{}, err + } + p.sqlDb.SetMaxOpenConns(10) + p.sqlDb.SetMaxIdleConns(10) + + err = p.sqlDb.Ping() + if err != nil { + return DatabaseProvider{}, err + } + + exists, err := p.tableExists("FileMetaData") + if err != nil { + return DatabaseProvider{}, err + } + if !exists { + return p, p.createNewDatabase() + } + return p, nil + } + return p, nil +} + +// Close the database connection +func (p DatabaseProvider) Close() { + if p.sqlDb != nil { + err := p.sqlDb.Close() + if err != nil { + fmt.Println(err) + } + } + p.sqlDb = nil +} + +// RunGarbageCollection runs the databases GC +func (p DatabaseProvider) RunGarbageCollection() { + p.cleanExpiredSessions() + p.cleanApiKeys() +} + +func (p DatabaseProvider) tableExists(tableName string) (bool, error) { + var count int + row := p.sqlDb.QueryRow("SELECT COUNT(*) FROM information_schema.tables WHERE table_schema = DATABASE() AND table_name = ?", tableName) + err := row.Scan(&count) + if err != nil { + return false, err + } + return count > 0, nil +} + +func (p DatabaseProvider) createNewDatabase() error { + statements := []string{ + `CREATE TABLE ApiKeys ( + Id VARCHAR(64) NOT NULL, + FriendlyName VARCHAR(255) NOT NULL, + LastUsed BIGINT NOT NULL, + Permissions INT NOT NULL DEFAULT 0, + Expiry BIGINT NOT NULL DEFAULT 0, + IsSystemKey TINYINT NOT NULL DEFAULT 0, + UserId INT NOT NULL, + PublicId VARCHAR(64) NOT NULL, + UploadRequestId VARCHAR(64) NOT NULL, + PRIMARY KEY (Id), + UNIQUE KEY idx_apikeys_publicid (PublicId) + )`, + `CREATE TABLE E2EConfig ( + Id INT NOT NULL AUTO_INCREMENT, + Config MEDIUMBLOB NOT NULL, + UserId INT NOT NULL, + PRIMARY KEY (Id), + UNIQUE KEY idx_e2econfig_userid (UserId) + )`, + `CREATE TABLE FileMetaData ( + Id VARCHAR(64) NOT NULL, + Name VARCHAR(1024) NOT NULL, + Size VARCHAR(64) NOT NULL, + SHA1 VARCHAR(64) NOT NULL, + ExpireAt BIGINT NOT NULL, + SizeBytes BIGINT NOT NULL, + DownloadsRemaining INT NOT NULL, + DownloadCount INT NOT NULL, + PasswordHash VARCHAR(255) NOT NULL, + HotlinkId VARCHAR(255) NOT NULL, + ContentType VARCHAR(255) NOT NULL, + AwsBucket VARCHAR(255) NOT NULL, + Encryption MEDIUMBLOB NOT NULL, + UnlimitedDownloads TINYINT NOT NULL, + UnlimitedTime TINYINT NOT NULL, + UserId INT NOT NULL, + UploadDate BIGINT NOT NULL, + PendingDeletion BIGINT NOT NULL, + UploadRequestId VARCHAR(64) NOT NULL, + PRIMARY KEY (Id) + )`, + `CREATE TABLE Hotlinks ( + Id VARCHAR(255) NOT NULL, + FileId VARCHAR(64) NOT NULL, + PRIMARY KEY (Id), + UNIQUE KEY idx_hotlinks_fileid (FileId) + )`, + `CREATE TABLE Sessions ( + Id VARCHAR(255) NOT NULL, + RenewAt BIGINT NOT NULL, + ValidUntil BIGINT NOT NULL, + UserId INT NOT NULL, + PRIMARY KEY (Id) + )`, + `CREATE TABLE Users ( + Id INT NOT NULL AUTO_INCREMENT, + Name VARCHAR(255) NOT NULL, + Password VARCHAR(255), + Permissions INT NOT NULL, + Userlevel INT NOT NULL, + LastOnline BIGINT NOT NULL DEFAULT 0, + ResetPassword TINYINT NOT NULL DEFAULT 0, + PRIMARY KEY (Id), + UNIQUE KEY idx_users_name (Name) + )`, + `CREATE TABLE UploadRequests ( + Id VARCHAR(64) NOT NULL, + Name VARCHAR(255) NOT NULL, + UserId INT NOT NULL, + Expiry BIGINT NOT NULL, + MaxFiles INT NOT NULL, + MaxSize INT NOT NULL, + Creation BIGINT NOT NULL, + ApiKey VARCHAR(255) NOT NULL, + Note TEXT NOT NULL, + PRIMARY KEY (Id), + UNIQUE KEY idx_uploadrequests_apikey (ApiKey) + )`, + `CREATE TABLE Statistics ( + Id INT NOT NULL AUTO_INCREMENT, + Type INT NOT NULL, + Value BIGINT, + PRIMARY KEY (Id), + UNIQUE KEY idx_statistics_type (Type) + )`, + `CREATE TABLE SchemaVersion ( + Version INT NOT NULL + )`, + } + for _, statement := range statements { + _, err := p.sqlDb.Exec(statement) + if err != nil { + return err + } + } + p.SetDbVersion(DatabaseSchemeVersion) + return nil +} diff --git a/internal/configuration/database/provider/mariadb/Mariadb_test.go b/internal/configuration/database/provider/mariadb/Mariadb_test.go new file mode 100644 index 00000000..b55ec658 --- /dev/null +++ b/internal/configuration/database/provider/mariadb/Mariadb_test.go @@ -0,0 +1,719 @@ +//go:build test && mariadbtest + +// Package mariadb tests require a real MariaDB/MySQL server. They are gated behind the +// "mariadbtest" build tag (not part of the default "test" tag set) and configured via +// environment variables, since there is no pure-Go in-process mock for MariaDB/MySQL +// (unlike SQLite, which is just a file, or Redis, which uses miniredis): +// +// GOKAPI_MARIADB_HOST=127.0.0.1:3306 +// GOKAPI_MARIADB_DBNAME=gokapi_test +// GOKAPI_MARIADB_USER=gokapi +// GOKAPI_MARIADB_PASSWORD=secret +// +// Run with: go test ./internal/configuration/database/provider/mariadb/... --tags=test,mariadbtest +// +// The target database must already exist; the tables in it are dropped and recreated on +// every run, so use a dedicated, disposable database - never point this at production data. +package mariadb + +import ( + "fmt" + "math" + "os" + "slices" + "sync" + "testing" + "time" + + "github.com/forceu/gokapi/internal/helper" + "github.com/forceu/gokapi/internal/models" + "github.com/forceu/gokapi/internal/test" +) + +var config = models.DbConnection{ + HostUrl: envOrDefault("GOKAPI_MARIADB_HOST", "127.0.0.1:3306"), + DatabaseName: envOrDefault("GOKAPI_MARIADB_DBNAME", "gokapi_test"), + Username: envOrDefault("GOKAPI_MARIADB_USER", "gokapi"), + Password: envOrDefault("GOKAPI_MARIADB_PASSWORD", "secret"), + Type: 2, // dbabstraction.TypeMariaDb +} + +func envOrDefault(key, fallback string) string { + if val, ok := os.LookupEnv(key); ok { + return val + } + return fallback +} + +var dbInstance DatabaseProvider + +func dropAllTables() error { + instance, err := New(config) + if err != nil { + return err + } + tables := []string{"ApiKeys", "E2EConfig", "FileMetaData", "Hotlinks", "Sessions", "Users", "UploadRequests", "Statistics", "SchemaVersion"} + for _, table := range tables { + _, err = instance.sqlDb.Exec("DROP TABLE IF EXISTS " + table) + if err != nil { + return err + } + } + instance.Close() + return nil +} + +func TestMain(m *testing.M) { + err := dropAllTables() + if err != nil { + fmt.Println("Could not reset MariaDB test database:", err) + os.Exit(1) + } + exitVal := m.Run() + os.Exit(exitVal) +} + +func TestInit(t *testing.T) { + instance, err := New(config) + test.IsNil(t, err) + instance.Close() + + _, err = New(models.DbConnection{HostUrl: "", DatabaseName: "gokapi_test", Type: 2}) + test.IsNotNil(t, err) + _, err = New(models.DbConnection{HostUrl: config.HostUrl, DatabaseName: "", Type: 2}) + test.IsNotNil(t, err) +} + +func TestClose(t *testing.T) { + instance, err := New(config) + test.IsNil(t, err) + instance.Close() + instance, err = New(config) + test.IsNil(t, err) + dbInstance = instance +} + +func TestDatabaseProvider_GetType(t *testing.T) { + test.IsEqualInt(t, dbInstance.GetType(), 2) +} + +func TestDatabaseProvider_GetDbVersion(t *testing.T) { + version := dbInstance.GetDbVersion() + test.IsEqualInt(t, version, DatabaseSchemeVersion) + dbInstance.SetDbVersion(99) + test.IsEqualInt(t, dbInstance.GetDbVersion(), 99) + dbInstance.SetDbVersion(DatabaseSchemeVersion) +} + +func TestDatabaseProvider_GetSchemaVersion(t *testing.T) { + test.IsEqualInt(t, dbInstance.GetSchemaVersion(), DatabaseSchemeVersion) +} + +func TestMetaData(t *testing.T) { + files := dbInstance.GetAllMetadata() + test.IsEqualInt(t, len(files), 0) + + dbInstance.SaveMetaData(models.File{Id: "testfile", Name: "test.txt", ExpireAt: time.Now().Add(time.Hour).Unix()}) + files = dbInstance.GetAllMetadata() + test.IsEqualInt(t, len(files), 1) + test.IsEqualString(t, files["testfile"].Name, "test.txt") + + file, ok := dbInstance.GetMetaDataById("testfile") + test.IsEqualBool(t, ok, true) + test.IsEqualString(t, file.Id, "testfile") + _, ok = dbInstance.GetMetaDataById("invalid") + test.IsEqualBool(t, ok, false) + + test.IsEqualInt(t, len(dbInstance.GetAllMetadata()), 1) + dbInstance.DeleteMetaData("invalid") + test.IsEqualInt(t, len(dbInstance.GetAllMetadata()), 1) + + test.IsEqualBool(t, file.UnlimitedDownloads, false) + test.IsEqualBool(t, file.UnlimitedTime, false) + + dbInstance.DeleteMetaData("testfile") + test.IsEqualInt(t, len(dbInstance.GetAllMetadata()), 0) + + dbInstance.SaveMetaData(models.File{ + Id: "test2", + Name: "test2", + UnlimitedDownloads: true, + UnlimitedTime: false, + }) + + file, ok = dbInstance.GetMetaDataById("test2") + test.IsEqualBool(t, ok, true) + test.IsEqualBool(t, file.UnlimitedDownloads, true) + test.IsEqualBool(t, file.UnlimitedTime, false) + + dbInstance.SaveMetaData(models.File{ + Id: "test3", + Name: "test3", + DownloadsRemaining: 4, + UnlimitedDownloads: false, + UnlimitedTime: true, + }) + file, ok = dbInstance.GetMetaDataById("test3") + test.IsEqualBool(t, ok, true) + test.IsEqualBool(t, file.UnlimitedDownloads, false) + test.IsEqualBool(t, file.UnlimitedTime, true) + test.IsEqualInt(t, file.DownloadsRemaining, 4) + remaining := dbInstance.GetDownloadsRemaining(file.Id) + test.IsEqualInt(t, remaining, 4) + + dbInstance.DeleteMetaData("test2") + dbInstance.DeleteMetaData("test3") +} + +func TestHotlink(t *testing.T) { + dbInstance.SaveHotlink(models.File{Id: "testfile", Name: "test.txt", HotlinkId: "testlink", ExpireAt: time.Now().Add(time.Hour).Unix()}) + + hotlink, ok := dbInstance.GetHotlink("testlink") + test.IsEqualBool(t, ok, true) + test.IsEqualString(t, hotlink, "testfile") + _, ok = dbInstance.GetHotlink("invalid") + test.IsEqualBool(t, ok, false) + + dbInstance.DeleteHotlink("invalid") + _, ok = dbInstance.GetHotlink("testlink") + test.IsEqualBool(t, ok, true) + dbInstance.DeleteHotlink("testlink") + _, ok = dbInstance.GetHotlink("testlink") + test.IsEqualBool(t, ok, false) + + dbInstance.SaveHotlink(models.File{Id: "testfile", Name: "test.txt", HotlinkId: "testlink", ExpireAt: 0, UnlimitedTime: true}) + hotlink, ok = dbInstance.GetHotlink("testlink") + test.IsEqualBool(t, ok, true) + test.IsEqualString(t, hotlink, "testfile") + + dbInstance.SaveHotlink(models.File{Id: "file2", Name: "file2.txt", HotlinkId: "link2", ExpireAt: time.Now().Add(time.Hour).Unix()}) + dbInstance.SaveHotlink(models.File{Id: "file3", Name: "file3.txt", HotlinkId: "link3", ExpireAt: time.Now().Add(time.Hour).Unix()}) + + hotlinks := dbInstance.GetAllHotlinks() + test.IsEqualInt(t, len(hotlinks), 3) + test.IsEqualBool(t, slices.Contains(hotlinks, "testlink"), true) + test.IsEqualBool(t, slices.Contains(hotlinks, "link2"), true) + test.IsEqualBool(t, slices.Contains(hotlinks, "link3"), true) + dbInstance.DeleteHotlink("") + hotlinks = dbInstance.GetAllHotlinks() + test.IsEqualInt(t, len(hotlinks), 3) + + dbInstance.DeleteHotlink("testlink") + dbInstance.DeleteHotlink("link2") + dbInstance.DeleteHotlink("link3") +} + +func TestDatabaseProvider_IncreaseDownloadCount(t *testing.T) { + newFile := models.File{ + Id: "newFileId", + Name: "newFileName", + Size: "3GB", + SHA1: "newSHA1", + PasswordHash: "newPassword", + HotlinkId: "newHotlink", + ContentType: "newContent", + AwsBucket: "newAws", + ExpireAt: 123456, + SizeBytes: 456789, + DownloadsRemaining: 11, + DownloadCount: 2, + Encryption: models.EncryptionInfo{ + IsEncrypted: true, + IsEndToEndEncrypted: true, + DecryptionKey: []byte("newDecryptionKey"), + Nonce: []byte("newDecryptionNonce"), + }, + UnlimitedDownloads: true, + UnlimitedTime: true, + } + dbInstance.SaveMetaData(newFile) + dbInstance.IncreaseDownloadCount(newFile.Id, false) + retrievedFile, ok := dbInstance.GetMetaDataById(newFile.Id) + test.IsEqualBool(t, ok, true) + test.IsEqualInt(t, retrievedFile.DownloadCount, 3) + test.IsEqualInt(t, retrievedFile.DownloadsRemaining, 11) + newFile.DownloadCount = 3 + test.IsEqual(t, retrievedFile, newFile) + + dbInstance.IncreaseDownloadCount(newFile.Id, true) + retrievedFile, ok = dbInstance.GetMetaDataById(newFile.Id) + test.IsEqualBool(t, ok, true) + test.IsEqualInt(t, retrievedFile.DownloadCount, 4) + test.IsEqualInt(t, retrievedFile.DownloadsRemaining, 10) + newFile.DownloadCount = 4 + newFile.DownloadsRemaining = 10 + test.IsEqual(t, retrievedFile, newFile) + dbInstance.DeleteMetaData(newFile.Id) +} + +func TestApiKey(t *testing.T) { + key1 := models.ApiKey{ + Id: "newkey", + FriendlyName: "New Key", + LastUsed: 100, + Permissions: 20, + PublicId: "_n3wkey", + Expiry: 0, + IsSystemKey: false, + UserId: 5, + } + key2 := models.ApiKey{ + Id: "newkey2", + FriendlyName: "New Key2", + PublicId: "_n3wkey2", + Expiry: 17362039396, + LastUsed: 200, + Permissions: 40, + IsSystemKey: true, + UserId: 10, + } + dbInstance.SaveApiKey(key1) + dbInstance.SaveApiKey(key2) + dbInstance.SaveApiKey(models.ApiKey{ + Id: "expiredKey", + PublicId: "expiredKey", + FriendlyName: "expiredKey", + Expiry: 1, + }) + + keys := dbInstance.GetAllApiKeys() + test.IsEqualInt(t, len(keys), 2) + test.IsEqual(t, keys["newkey"], key1) + test.IsEqual(t, keys["newkey2"], key2) + dbInstance.DeleteApiKey("newkey2") + test.IsEqualInt(t, len(dbInstance.GetAllApiKeys()), 1) + + key, ok := dbInstance.GetApiKey("newkey") + test.IsEqualBool(t, ok, true) + test.IsEqual(t, key, key1) + _, ok = dbInstance.GetApiKey("newkey2") + test.IsEqualBool(t, ok, false) + + dbInstance.SaveApiKey(models.ApiKey{ + Id: "newkey", + FriendlyName: "Old Key", + LastUsed: 100, + }) + key, ok = dbInstance.GetApiKey("newkey") + test.IsEqualBool(t, ok, true) + test.IsEqualString(t, key.FriendlyName, "Old Key") + + dbInstance.DeleteApiKey("newkey") + dbInstance.DeleteApiKey("expiredKey") +} + +func TestSession(t *testing.T) { + renewAt := time.Now().Add(1 * time.Hour).Unix() + dbInstance.SaveSession("newsession", models.Session{ + RenewAt: renewAt, + ValidUntil: time.Now().Add(2 * time.Hour).Unix(), + }) + + session, ok := dbInstance.GetSession("newsession") + test.IsEqualBool(t, ok, true) + test.IsEqualBool(t, session.RenewAt == renewAt, true) + + dbInstance.DeleteSession("newsession") + _, ok = dbInstance.GetSession("newsession") + test.IsEqualBool(t, ok, false) + + dbInstance.SaveSession("newsession", models.Session{ + RenewAt: renewAt, + ValidUntil: time.Now().Add(2 * time.Hour).Unix(), + }) + + dbInstance.SaveSession("anothersession", models.Session{ + RenewAt: renewAt, + ValidUntil: time.Now().Add(2 * time.Hour).Unix(), + }) + _, ok = dbInstance.GetSession("newsession") + test.IsEqualBool(t, ok, true) + _, ok = dbInstance.GetSession("anothersession") + test.IsEqualBool(t, ok, true) + + dbInstance.DeleteAllSessions() + _, ok = dbInstance.GetSession("newsession") + test.IsEqualBool(t, ok, false) + _, ok = dbInstance.GetSession("anothersession") + test.IsEqualBool(t, ok, false) + + session = models.Session{ + RenewAt: 2147483645, + ValidUntil: 2147483645, + UserId: 20, + } + dbInstance.SaveSession("sess_user1", session) + dbInstance.SaveSession("sess_user2", session) + dbInstance.SaveSession("sess_user3", session) + session.UserId = 40 + dbInstance.SaveSession("sess_user4", session) + _, ok = dbInstance.GetSession("sess_user1") + test.IsEqualBool(t, ok, true) + _, ok = dbInstance.GetSession("sess_user2") + test.IsEqualBool(t, ok, true) + _, ok = dbInstance.GetSession("sess_user3") + test.IsEqualBool(t, ok, true) + _, ok = dbInstance.GetSession("sess_user4") + test.IsEqualBool(t, ok, true) + dbInstance.DeleteAllSessionsByUser(20) + _, ok = dbInstance.GetSession("sess_user1") + test.IsEqualBool(t, ok, false) + _, ok = dbInstance.GetSession("sess_user2") + test.IsEqualBool(t, ok, false) + _, ok = dbInstance.GetSession("sess_user3") + test.IsEqualBool(t, ok, false) + _, ok = dbInstance.GetSession("sess_user4") + test.IsEqualBool(t, ok, true) + dbInstance.DeleteAllSessions() +} + +func TestFileRequest(t *testing.T) { + req1 := models.FileRequest{ + Id: "req1", + Name: "New file request", + UserId: 45564, + ApiKey: "123", + CreationDate: time.Now().Unix(), + } + dbInstance.SaveFileRequest(req1) + + request, ok := dbInstance.GetFileRequest("req1") + test.IsEqualBool(t, ok, true) + test.IsEqualString(t, request.Id, "req1") + test.IsEqualString(t, request.Name, "New file request") + + _, ok = dbInstance.GetFileRequest("invalid") + test.IsEqualBool(t, ok, false) + + _, ok = dbInstance.GetFileRequest("") + test.IsEqualBool(t, ok, false) + + dbInstance.DeleteFileRequest(models.FileRequest{Id: "invalid"}) + _, ok = dbInstance.GetFileRequest("req1") + test.IsEqualBool(t, ok, true) + + dbInstance.DeleteFileRequest(req1) + _, ok = dbInstance.GetFileRequest("req1") + test.IsEqualBool(t, ok, false) + + req2 := models.FileRequest{ + Id: "req2", + UserId: 45564, + Name: "file2.txt", + ApiKey: "456", + CreationDate: time.Now().Add(-time.Minute).Unix(), + } + req3 := models.FileRequest{ + Id: "req3", + Name: "file3.txt", + UserId: 45564, + ApiKey: "789", + CreationDate: time.Now().Add(-2 * time.Minute).Unix(), + } + + dbInstance.SaveFileRequest(req1) + dbInstance.SaveFileRequest(req2) + dbInstance.SaveFileRequest(req3) + + requests := dbInstance.GetAllFileRequests() + test.IsEqualInt(t, len(requests), 3) + + ids := []string{requests[0].Id, requests[1].Id, requests[2].Id} + test.IsEqualBool(t, slices.Contains(ids, "req1"), true) + test.IsEqualBool(t, slices.Contains(ids, "req2"), true) + test.IsEqualBool(t, slices.Contains(ids, "req3"), true) + + test.IsEqualBool(t, requests[0].CreationDate >= requests[1].CreationDate, true) + test.IsEqualBool(t, requests[1].CreationDate >= requests[2].CreationDate, true) + + dbInstance.DeleteFileRequest(req1) + dbInstance.DeleteFileRequest(req2) + dbInstance.DeleteFileRequest(req3) +} + +func TestGarbageCollectionSessions(t *testing.T) { + dbInstance.SaveSession("todelete1", models.Session{ + RenewAt: time.Now().Add(-10 * time.Second).Unix(), + ValidUntil: time.Now().Add(-10 * time.Second).Unix(), + }) + dbInstance.SaveSession("todelete2", models.Session{ + RenewAt: time.Now().Add(10 * time.Second).Unix(), + ValidUntil: time.Now().Add(-10 * time.Second).Unix(), + }) + dbInstance.SaveSession("tokeep1", models.Session{ + RenewAt: time.Now().Add(-10 * time.Second).Unix(), + ValidUntil: time.Now().Add(10 * time.Second).Unix(), + }) + dbInstance.SaveSession("tokeep2", models.Session{ + RenewAt: time.Now().Add(10 * time.Second).Unix(), + ValidUntil: time.Now().Add(10 * time.Second).Unix(), + }) + for _, item := range []string{"todelete1", "todelete2", "tokeep1", "tokeep2"} { + _, result := dbInstance.GetSession(item) + test.IsEqualBool(t, result, true) + } + dbInstance.RunGarbageCollection() + for _, item := range []string{"todelete1", "todelete2"} { + _, result := dbInstance.GetSession(item) + test.IsEqualBool(t, result, false) + } + for _, item := range []string{"tokeep1", "tokeep2"} { + _, result := dbInstance.GetSession(item) + test.IsEqualBool(t, result, true) + } + dbInstance.DeleteAllSessions() +} + +func TestEnd2EndInfo(t *testing.T) { + info := dbInstance.GetEnd2EndInfo(4) + test.IsEqualInt(t, info.Version, 0) + test.IsEqualBool(t, info.HasBeenSetUp(), false) + + dbInstance.SaveEnd2EndInfo(models.E2EInfoEncrypted{ + Version: 1, + Nonce: []byte("testNonce1"), + Content: []byte("testContent1"), + AvailableFiles: nil, + }, 4) + + info = dbInstance.GetEnd2EndInfo(4) + test.IsEqualInt(t, info.Version, 1) + test.IsEqualBool(t, info.HasBeenSetUp(), true) + test.IsEqualByteSlice(t, info.Nonce, []byte("testNonce1")) + test.IsEqualByteSlice(t, info.Content, []byte("testContent1")) + test.IsEqualBool(t, len(info.AvailableFiles) == 0, true) + + dbInstance.SaveEnd2EndInfo(models.E2EInfoEncrypted{ + Version: 2, + Nonce: []byte("testNonce2"), + Content: []byte("testContent2"), + AvailableFiles: nil, + }, 4) + + info = dbInstance.GetEnd2EndInfo(4) + test.IsEqualInt(t, info.Version, 2) + test.IsEqualBool(t, info.HasBeenSetUp(), true) + test.IsEqualByteSlice(t, info.Nonce, []byte("testNonce2")) + test.IsEqualByteSlice(t, info.Content, []byte("testContent2")) + test.IsEqualBool(t, len(info.AvailableFiles) == 0, true) + + dbInstance.DeleteEnd2EndInfo(4) + info = dbInstance.GetEnd2EndInfo(4) + test.IsEqualInt(t, info.Version, 0) + test.IsEqualBool(t, info.HasBeenSetUp(), false) +} + +func TestUpdateTimeApiKey(t *testing.T) { + retrievedKey, ok := dbInstance.GetApiKey("key1") + test.IsEqualBool(t, ok, false) + test.IsEqualString(t, retrievedKey.Id, "") + + key := models.ApiKey{ + Id: "key1", + FriendlyName: "key1", + PublicId: "key1", + LastUsed: 100, + } + dbInstance.SaveApiKey(key) + key = models.ApiKey{ + Id: "key2", + FriendlyName: "key2", + PublicId: "key2", + LastUsed: 200, + } + dbInstance.SaveApiKey(key) + + retrievedKey, ok = dbInstance.GetApiKey("key1") + test.IsEqualBool(t, ok, true) + test.IsEqualString(t, retrievedKey.Id, "key1") + test.IsEqualInt64(t, retrievedKey.LastUsed, 100) + retrievedKey, ok = dbInstance.GetApiKey("key2") + test.IsEqualBool(t, ok, true) + test.IsEqualString(t, retrievedKey.Id, "key2") + test.IsEqualInt64(t, retrievedKey.LastUsed, 200) + + key.LastUsed = 300 + dbInstance.UpdateTimeApiKey(key) + + retrievedKey, ok = dbInstance.GetApiKey("key1") + test.IsEqualBool(t, ok, true) + test.IsEqualString(t, retrievedKey.Id, "key1") + test.IsEqualInt64(t, retrievedKey.LastUsed, 100) + retrievedKey, ok = dbInstance.GetApiKey("key2") + test.IsEqualBool(t, ok, true) + test.IsEqualString(t, retrievedKey.Id, "key2") + test.IsEqualInt64(t, retrievedKey.LastUsed, 300) + + dbInstance.SaveApiKey(models.ApiKey{ + Id: "publicTest", + PublicId: "publicId", + }) + _, ok = dbInstance.GetApiKey("publicTest") + test.IsEqualBool(t, ok, true) + _, ok = dbInstance.GetApiKey("publicId") + test.IsEqualBool(t, ok, false) + _, ok = dbInstance.GetApiKeyByPublicKey("publicTest") + test.IsEqualBool(t, ok, false) + keyName, ok := dbInstance.GetApiKeyByPublicKey("publicId") + test.IsEqualBool(t, ok, true) + test.IsEqualString(t, keyName, "publicTest") + + dbInstance.DeleteApiKey("key1") + dbInstance.DeleteApiKey("key2") + dbInstance.DeleteApiKey("publicTest") +} + +func TestParallelConnectionsWritingAndReading(t *testing.T) { + var wg sync.WaitGroup + + simulatedConnection := func(t *testing.T) { + file := models.File{ + Id: helper.GenerateRandomString(10), + Name: helper.GenerateRandomString(10), + Size: "10B", + SHA1: "1289423794287598237489", + ExpireAt: math.MaxInt32, + SizeBytes: 10, + DownloadsRemaining: 10, + DownloadCount: 10, + Encryption: models.EncryptionInfo{}, + UnlimitedDownloads: false, + UnlimitedTime: false, + } + dbInstance.SaveMetaData(file) + retrievedFile, ok := dbInstance.GetMetaDataById(file.Id) + test.IsEqualBool(t, ok, true) + test.IsEqualString(t, retrievedFile.Name, file.Name) + dbInstance.DeleteMetaData(file.Id) + _, ok = dbInstance.GetMetaDataById(file.Id) + test.IsEqualBool(t, ok, false) + } + + for i := 1; i <= 100; i++ { + wg.Add(1) + go func() { + defer wg.Done() + simulatedConnection(t) + }() + } + wg.Wait() +} + +func TestParallelConnectionsReading(t *testing.T) { + var wg sync.WaitGroup + + dbInstance.SaveApiKey(models.ApiKey{ + Id: "readtest", + FriendlyName: "readtest", + LastUsed: 40000, + }) + simulatedConnection := func(t *testing.T) { + _, ok := dbInstance.GetApiKey("readtest") + test.IsEqualBool(t, ok, true) + } + + for i := 1; i <= 1000; i++ { + wg.Add(1) + go func() { + defer wg.Done() + simulatedConnection(t) + }() + } + wg.Wait() + dbInstance.DeleteApiKey("readtest") +} + +func TestStatistics(t *testing.T) { + test.IsEqualInt64(t, int64(dbInstance.GetStatTraffic()), 0) + dbInstance.SaveStatTraffic(1024) + test.IsEqualInt64(t, int64(dbInstance.GetStatTraffic()), 1024) + dbInstance.SaveStatTraffic(2048) + test.IsEqualInt64(t, int64(dbInstance.GetStatTraffic()), 2048) + + _, ok := dbInstance.GetTrafficSince() + test.IsEqualBool(t, ok, false) + dbInstance.SaveTrafficSince(12345) + since, ok := dbInstance.GetTrafficSince() + test.IsEqualBool(t, ok, true) + test.IsEqualInt64(t, since, 12345) + dbInstance.SaveTrafficSince(54321) + since, ok = dbInstance.GetTrafficSince() + test.IsEqualBool(t, ok, true) + test.IsEqualInt64(t, since, 54321) +} + +func TestUsers(t *testing.T) { + users := dbInstance.GetAllUsers() + test.IsEqualInt(t, len(users), 0) + user := models.User{ + Id: 2, + Name: "test", + Permissions: models.UserPermissionAll, + UserLevel: models.UserLevelUser, + LastOnline: 1337, + Password: "123456", + ResetPassword: true, + } + dbInstance.SaveUser(user, false) + retrievedUser, ok := dbInstance.GetUser(2) + test.IsEqualBool(t, ok, true) + test.IsEqual(t, retrievedUser, user) + users = dbInstance.GetAllUsers() + test.IsEqualInt(t, len(users), 1) + test.IsEqualInt(t, retrievedUser.Id, 2) + + _, ok = dbInstance.GetUser(0) + test.IsEqualBool(t, ok, false) + _, ok = dbInstance.GetUserByName("invalid") + test.IsEqualBool(t, ok, false) + retrievedUser, ok = dbInstance.GetUserByName("test") + test.IsEqualBool(t, ok, true) + test.IsEqual(t, retrievedUser, user) + + dbInstance.DeleteUser(2) + _, ok = dbInstance.GetUser(2) + test.IsEqualBool(t, ok, false) + + user = models.User{ + Id: 1000, + Name: "test2", + Permissions: models.UserPermissionNone, + UserLevel: models.UserLevelAdmin, + LastOnline: 1338, + Password: "1234568", + ResetPassword: true, + } + dbInstance.SaveUser(user, true) + _, ok = dbInstance.GetUser(1000) + test.IsEqualBool(t, ok, false) + retrievedUser, ok = dbInstance.GetUserByName("test2") + test.IsEqualBool(t, ok, true) + test.IsEqualBool(t, retrievedUser.Id == 1000, false) + user.Id = retrievedUser.Id + test.IsEqual(t, retrievedUser, user) + + dbInstance.UpdateUserLastOnline(retrievedUser.Id) + retrievedUser, ok = dbInstance.GetUser(retrievedUser.Id) + test.IsEqualBool(t, ok, true) + test.IsEqualBool(t, time.Now().Unix()-retrievedUser.LastOnline < 5, true) + test.IsEqualBool(t, time.Now().Unix()-retrievedUser.LastOnline > -1, true) + + user.Name = "test1" + dbInstance.SaveUser(user, true) + user.Name = "test3" + dbInstance.SaveUser(user, true) + user.Name = "test99" + user.UserLevel = models.UserLevelSuperAdmin + dbInstance.SaveUser(user, true) + user.Name = "test0" + user.UserLevel = models.UserLevelUser + dbInstance.SaveUser(user, true) + + users = dbInstance.GetAllUsers() + test.IsEqualInt(t, len(users), 5) + test.IsEqualString(t, users[0].Name, "test99") + test.IsEqualString(t, users[1].Name, "test2") + test.IsEqualString(t, users[2].Name, "test1") + test.IsEqualString(t, users[3].Name, "test3") + test.IsEqualString(t, users[4].Name, "test0") +} diff --git a/internal/configuration/database/provider/mariadb/apikeys.go b/internal/configuration/database/provider/mariadb/apikeys.go new file mode 100644 index 00000000..395b71b2 --- /dev/null +++ b/internal/configuration/database/provider/mariadb/apikeys.go @@ -0,0 +1,127 @@ +package mariadb + +import ( + "database/sql" + "errors" + "time" + + "github.com/forceu/gokapi/internal/helper" + "github.com/forceu/gokapi/internal/models" +) + +type schemaApiKeys struct { + Id string + FriendlyName string + LastUsed int64 + Permissions int + Expiry int64 + IsSystemKey int + UserId int + PublicId string + UploadRequestId string +} + +// currentTime is used in order to modify the current time for testing purposes in unit tests +var currentTime = func() time.Time { + return time.Now() +} + +// GetAllApiKeys returns a map with all API keys +func (p DatabaseProvider) GetAllApiKeys() map[string]models.ApiKey { + result := make(map[string]models.ApiKey) + + rows, err := p.sqlDb.Query("SELECT * FROM ApiKeys WHERE ApiKeys.Expiry = 0 OR ApiKeys.Expiry > ?", currentTime().Unix()) + helper.Check(err) + defer rows.Close() + for rows.Next() { + rowData := schemaApiKeys{} + err = rows.Scan(&rowData.Id, &rowData.FriendlyName, &rowData.LastUsed, &rowData.Permissions, &rowData.Expiry, + &rowData.IsSystemKey, &rowData.UserId, &rowData.PublicId, &rowData.UploadRequestId) + helper.Check(err) + result[rowData.Id] = models.ApiKey{ + Id: rowData.Id, + PublicId: rowData.PublicId, + FriendlyName: rowData.FriendlyName, + LastUsed: rowData.LastUsed, + Permissions: models.ApiPermission(rowData.Permissions), + Expiry: rowData.Expiry, + IsSystemKey: rowData.IsSystemKey == 1, + UserId: rowData.UserId, + UploadRequestId: rowData.UploadRequestId, + } + } + return result +} + +// GetApiKey returns a models.ApiKey if valid or false if the ID is not valid +func (p DatabaseProvider) GetApiKey(id string) (models.ApiKey, bool) { + var rowResult schemaApiKeys + row := p.sqlDb.QueryRow("SELECT * FROM ApiKeys WHERE Id = ?", id) + err := row.Scan(&rowResult.Id, &rowResult.FriendlyName, &rowResult.LastUsed, &rowResult.Permissions, &rowResult.Expiry, + &rowResult.IsSystemKey, &rowResult.UserId, &rowResult.PublicId, &rowResult.UploadRequestId) + if err != nil { + if errors.Is(err, sql.ErrNoRows) { + return models.ApiKey{}, false + } + helper.Check(err) + return models.ApiKey{}, false + } + + result := models.ApiKey{ + Id: rowResult.Id, + PublicId: rowResult.PublicId, + FriendlyName: rowResult.FriendlyName, + LastUsed: rowResult.LastUsed, + Permissions: models.ApiPermission(rowResult.Permissions), + Expiry: rowResult.Expiry, + IsSystemKey: rowResult.IsSystemKey == 1, + UserId: rowResult.UserId, + UploadRequestId: rowResult.UploadRequestId, + } + + return result, true +} + +// GetApiKeyByPublicKey returns an API key by using the public key +func (p DatabaseProvider) GetApiKeyByPublicKey(publicKey string) (string, bool) { + var rowResult schemaApiKeys + row := p.sqlDb.QueryRow("SELECT Id FROM ApiKeys WHERE PublicId = ? LIMIT 1", publicKey) + err := row.Scan(&rowResult.Id) + if err != nil { + if errors.Is(err, sql.ErrNoRows) { + return "", false + } + helper.Check(err) + return "", false + } + return rowResult.Id, true +} + +// SaveApiKey saves the API key to the database +func (p DatabaseProvider) SaveApiKey(apikey models.ApiKey) { + isSystemKey := 0 + if apikey.IsSystemKey { + isSystemKey = 1 + } + _, err := p.sqlDb.Exec("REPLACE INTO ApiKeys (Id, FriendlyName, LastUsed, Permissions, Expiry, IsSystemKey, UserId, PublicId, UploadRequestId) VALUES (?, ?, ?, ?, ?, ?, ?, ?, ?)", + apikey.Id, apikey.FriendlyName, apikey.LastUsed, apikey.Permissions, apikey.Expiry, isSystemKey, apikey.UserId, apikey.PublicId, apikey.UploadRequestId) + helper.Check(err) +} + +// UpdateTimeApiKey writes the content of LastUsage to the database +func (p DatabaseProvider) UpdateTimeApiKey(apikey models.ApiKey) { + _, err := p.sqlDb.Exec("UPDATE ApiKeys SET LastUsed = ? WHERE Id = ?", + apikey.LastUsed, apikey.Id) + helper.Check(err) +} + +// DeleteApiKey deletes an API key with the given ID +func (p DatabaseProvider) DeleteApiKey(id string) { + _, err := p.sqlDb.Exec("DELETE FROM ApiKeys WHERE Id = ?", id) + helper.Check(err) +} + +func (p DatabaseProvider) cleanApiKeys() { + _, err := p.sqlDb.Exec("DELETE FROM ApiKeys WHERE ApiKeys.Expiry > 0 AND ApiKeys.Expiry < ?", currentTime().Unix()) + helper.Check(err) +} diff --git a/internal/configuration/database/provider/mariadb/e2econfig.go b/internal/configuration/database/provider/mariadb/e2econfig.go new file mode 100644 index 00000000..66a62f76 --- /dev/null +++ b/internal/configuration/database/provider/mariadb/e2econfig.go @@ -0,0 +1,57 @@ +package mariadb + +import ( + "bytes" + "database/sql" + "encoding/gob" + "errors" + + "github.com/forceu/gokapi/internal/helper" + "github.com/forceu/gokapi/internal/models" +) + +type schemaE2EConfig struct { + Id int64 + Config []byte + UserId int +} + +// SaveEnd2EndInfo stores the encrypted e2e info +func (p DatabaseProvider) SaveEnd2EndInfo(info models.E2EInfoEncrypted, userId int) { + var buf bytes.Buffer + enc := gob.NewEncoder(&buf) + err := enc.Encode(info) + helper.Check(err) + + _, err = p.sqlDb.Exec("INSERT INTO E2EConfig (Config, UserId) VALUES (?, ?) ON DUPLICATE KEY UPDATE Config = ?", + buf.Bytes(), userId, buf.Bytes()) + helper.Check(err) +} + +// GetEnd2EndInfo retrieves the encrypted e2e info +func (p DatabaseProvider) GetEnd2EndInfo(userId int) models.E2EInfoEncrypted { + result := models.E2EInfoEncrypted{} + rowResult := schemaE2EConfig{} + + row := p.sqlDb.QueryRow("SELECT Config FROM E2EConfig WHERE UserId = ?", userId) + err := row.Scan(&rowResult.Config) + if err != nil { + if errors.Is(err, sql.ErrNoRows) { + return result + } + helper.Check(err) + return result + } + + buf := bytes.NewBuffer(rowResult.Config) + dec := gob.NewDecoder(buf) + err = dec.Decode(&result) + helper.Check(err) + return result +} + +// DeleteEnd2EndInfo resets the encrypted e2e info +func (p DatabaseProvider) DeleteEnd2EndInfo(userId int) { + _, err := p.sqlDb.Exec("DELETE FROM E2EConfig WHERE UserId = ?", userId) + helper.Check(err) +} diff --git a/internal/configuration/database/provider/mariadb/filerequests.go b/internal/configuration/database/provider/mariadb/filerequests.go new file mode 100644 index 00000000..05a21b05 --- /dev/null +++ b/internal/configuration/database/provider/mariadb/filerequests.go @@ -0,0 +1,107 @@ +package mariadb + +import ( + "database/sql" + "errors" + + "github.com/forceu/gokapi/internal/helper" + "github.com/forceu/gokapi/internal/models" +) + +type schemaFileRequests struct { + Id string + Name string + UserId int + Expiry int64 + MaxFiles int + MaxSize int + Creation int64 + ApiKey string + Note string +} + +// GetFileRequest returns the FileRequest or false if not found +func (p DatabaseProvider) GetFileRequest(id string) (models.FileRequest, bool) { + if id == "" { + return models.FileRequest{}, false + } + var rowResult schemaFileRequests + row := p.sqlDb.QueryRow("SELECT * FROM UploadRequests WHERE Id = ?", id) + err := row.Scan(&rowResult.Id, &rowResult.Name, &rowResult.UserId, &rowResult.Expiry, + &rowResult.MaxFiles, &rowResult.MaxSize, &rowResult.Creation, &rowResult.ApiKey, &rowResult.Note) + if err != nil { + if errors.Is(err, sql.ErrNoRows) { + return models.FileRequest{}, false + } + helper.Check(err) + return models.FileRequest{}, false + } + result := models.FileRequest{ + Id: rowResult.Id, + Name: rowResult.Name, + UserId: rowResult.UserId, + MaxFiles: rowResult.MaxFiles, + MaxSize: rowResult.MaxSize, + Expiry: rowResult.Expiry, + CreationDate: rowResult.Creation, + ApiKey: rowResult.ApiKey, + Notes: rowResult.Note, + } + return result, true +} + +// GetAllFileRequests returns an array with all file requests, ordered by creation date +func (p DatabaseProvider) GetAllFileRequests() []models.FileRequest { + result := make([]models.FileRequest, 0) + rows, err := p.sqlDb.Query("SELECT * FROM UploadRequests ORDER BY Creation DESC, Name") + helper.Check(err) + defer rows.Close() + for rows.Next() { + rowData := schemaFileRequests{} + err = rows.Scan(&rowData.Id, &rowData.Name, &rowData.UserId, &rowData.Expiry, &rowData.MaxFiles, + &rowData.MaxSize, &rowData.Creation, &rowData.ApiKey, &rowData.Note) + helper.Check(err) + result = append(result, models.FileRequest{ + Id: rowData.Id, + Name: rowData.Name, + UserId: rowData.UserId, + MaxFiles: rowData.MaxFiles, + MaxSize: rowData.MaxSize, + Expiry: rowData.Expiry, + CreationDate: rowData.Creation, + ApiKey: rowData.ApiKey, + Notes: rowData.Note, + }) + } + return result +} + +// SaveFileRequest stores the file request associated with the file in the database +func (p DatabaseProvider) SaveFileRequest(request models.FileRequest) { + newData := schemaFileRequests{ + Id: request.Id, + Name: request.Name, + UserId: request.UserId, + MaxFiles: request.MaxFiles, + MaxSize: request.MaxSize, + Expiry: request.Expiry, + Creation: request.CreationDate, + ApiKey: request.ApiKey, + Note: request.Notes, + } + + _, err := p.sqlDb.Exec(`REPLACE INTO UploadRequests + (Id, Name, UserId, Expiry, MaxFiles, MaxSize, Creation, ApiKey, Note) + VALUES (?, ?, ?, ?, ?, ?, ?, ?, ?)`, + newData.Id, newData.Name, newData.UserId, newData.Expiry, newData.MaxFiles, newData.MaxSize, newData.Creation, newData.ApiKey, newData.Note) + helper.Check(err) +} + +// DeleteFileRequest deletes a file request with the given ID +func (p DatabaseProvider) DeleteFileRequest(request models.FileRequest) { + if request.Id == "" { + return + } + _, err := p.sqlDb.Exec("DELETE FROM UploadRequests WHERE Id = ?", request.Id) + helper.Check(err) +} diff --git a/internal/configuration/database/provider/mariadb/hotlinks.go b/internal/configuration/database/provider/mariadb/hotlinks.go new file mode 100644 index 00000000..84cbf273 --- /dev/null +++ b/internal/configuration/database/provider/mariadb/hotlinks.go @@ -0,0 +1,65 @@ +package mariadb + +import ( + "database/sql" + "errors" + + "github.com/forceu/gokapi/internal/helper" + "github.com/forceu/gokapi/internal/models" +) + +type schemaHotlinks struct { + Id string + FileId string +} + +// GetHotlink returns the id of the file associated or false if not found +func (p DatabaseProvider) GetHotlink(id string) (string, bool) { + var rowResult schemaHotlinks + row := p.sqlDb.QueryRow("SELECT FileId FROM Hotlinks WHERE Id = ?", id) + err := row.Scan(&rowResult.FileId) + if err != nil { + if errors.Is(err, sql.ErrNoRows) { + return "", false + } + helper.Check(err) + return "", false + } + return rowResult.FileId, true +} + +// GetAllHotlinks returns an array with all hotlink ids +func (p DatabaseProvider) GetAllHotlinks() []string { + ids := make([]string, 0) + rows, err := p.sqlDb.Query("SELECT Id FROM Hotlinks") + helper.Check(err) + defer rows.Close() + for rows.Next() { + rowData := schemaHotlinks{} + err = rows.Scan(&rowData.Id) + helper.Check(err) + ids = append(ids, rowData.Id) + } + return ids +} + +// SaveHotlink stores the hotlink associated with the file in the database +func (p DatabaseProvider) SaveHotlink(file models.File) { + newData := schemaHotlinks{ + Id: file.HotlinkId, + FileId: file.Id, + } + + _, err := p.sqlDb.Exec("REPLACE INTO Hotlinks (Id, FileId) VALUES (?, ?)", + newData.Id, newData.FileId) + helper.Check(err) +} + +// DeleteHotlink deletes a hotlink with the given hotlink ID +func (p DatabaseProvider) DeleteHotlink(id string) { + if id == "" { + return + } + _, err := p.sqlDb.Exec("DELETE FROM Hotlinks WHERE Id = ?", id) + helper.Check(err) +} diff --git a/internal/configuration/database/provider/mariadb/metadata.go b/internal/configuration/database/provider/mariadb/metadata.go new file mode 100644 index 00000000..fe08dfb0 --- /dev/null +++ b/internal/configuration/database/provider/mariadb/metadata.go @@ -0,0 +1,184 @@ +package mariadb + +import ( + "bytes" + "database/sql" + "encoding/gob" + "errors" + + "github.com/forceu/gokapi/internal/helper" + "github.com/forceu/gokapi/internal/models" +) + +type schemaMetaData struct { + Id string + Name string + Size string + SHA1 string + ExpireAt int64 + SizeBytes int64 + DownloadsRemaining int + DownloadCount int + PasswordHash string + HotlinkId string + ContentType string + AwsBucket string + Encryption []byte + UnlimitedDownloads int + UnlimitedTime int + UserId int + UploadDate int64 + PendingDeletion int64 + UploadRequestId string +} + +func (rowData schemaMetaData) ToFileModel() (models.File, error) { + result := models.File{ + Id: rowData.Id, + Name: rowData.Name, + Size: rowData.Size, + SHA1: rowData.SHA1, + ExpireAt: rowData.ExpireAt, + SizeBytes: rowData.SizeBytes, + DownloadsRemaining: rowData.DownloadsRemaining, + DownloadCount: rowData.DownloadCount, + PasswordHash: rowData.PasswordHash, + HotlinkId: rowData.HotlinkId, + ContentType: rowData.ContentType, + AwsBucket: rowData.AwsBucket, + Encryption: models.EncryptionInfo{}, + UnlimitedDownloads: rowData.UnlimitedDownloads == 1, + UnlimitedTime: rowData.UnlimitedTime == 1, + UserId: rowData.UserId, + UploadDate: rowData.UploadDate, + PendingDeletion: rowData.PendingDeletion, + UploadRequestId: rowData.UploadRequestId, + } + + buf := bytes.NewBuffer(rowData.Encryption) + dec := gob.NewDecoder(buf) + err := dec.Decode(&result.Encryption) + return result, err +} + +// GetAllMetadata returns a map of all available files +func (p DatabaseProvider) GetAllMetadata() map[string]models.File { + result := make(map[string]models.File) + rows, err := p.sqlDb.Query("SELECT * FROM FileMetaData") + helper.Check(err) + defer rows.Close() + for rows.Next() { + rowData := schemaMetaData{} + err = rows.Scan(&rowData.Id, &rowData.Name, &rowData.Size, &rowData.SHA1, &rowData.ExpireAt, &rowData.SizeBytes, + &rowData.DownloadsRemaining, &rowData.DownloadCount, &rowData.PasswordHash, &rowData.HotlinkId, &rowData.ContentType, + &rowData.AwsBucket, &rowData.Encryption, &rowData.UnlimitedDownloads, &rowData.UnlimitedTime, &rowData.UserId, + &rowData.UploadDate, &rowData.PendingDeletion, &rowData.UploadRequestId) + helper.Check(err) + var metaData models.File + metaData, err = rowData.ToFileModel() + helper.Check(err) + result[metaData.Id] = metaData + } + return result +} + +// GetMetaDataById returns a models.File from the ID passed or false if the id is not valid +func (p DatabaseProvider) GetMetaDataById(id string) (models.File, bool) { + result := models.File{} + rowData := schemaMetaData{} + + row := p.sqlDb.QueryRow("SELECT * FROM FileMetaData WHERE Id = ?", id) + err := row.Scan(&rowData.Id, &rowData.Name, &rowData.Size, &rowData.SHA1, &rowData.ExpireAt, &rowData.SizeBytes, + &rowData.DownloadsRemaining, &rowData.DownloadCount, &rowData.PasswordHash, + &rowData.HotlinkId, &rowData.ContentType, &rowData.AwsBucket, &rowData.Encryption, + &rowData.UnlimitedDownloads, &rowData.UnlimitedTime, &rowData.UserId, &rowData.UploadDate, + &rowData.PendingDeletion, &rowData.UploadRequestId) + if err != nil { + if errors.Is(err, sql.ErrNoRows) { + return result, false + } + helper.Check(err) + return result, false + } + result, err = rowData.ToFileModel() + helper.Check(err) + return result, true +} + +// SaveMetaData stores the metadata of a file to the disk +func (p DatabaseProvider) SaveMetaData(file models.File) { + newData := schemaMetaData{ + Id: file.Id, + Name: file.Name, + Size: file.Size, + SHA1: file.SHA1, + ExpireAt: file.ExpireAt, + SizeBytes: file.SizeBytes, + DownloadsRemaining: file.DownloadsRemaining, + DownloadCount: file.DownloadCount, + PasswordHash: file.PasswordHash, + HotlinkId: file.HotlinkId, + ContentType: file.ContentType, + AwsBucket: file.AwsBucket, + UserId: file.UserId, + UploadDate: file.UploadDate, + PendingDeletion: file.PendingDeletion, + UploadRequestId: file.UploadRequestId, + } + + if file.UnlimitedDownloads { + newData.UnlimitedDownloads = 1 + } + if file.UnlimitedTime { + newData.UnlimitedTime = 1 + } + + var buf bytes.Buffer + enc := gob.NewEncoder(&buf) + err := enc.Encode(file.Encryption) + helper.Check(err) + newData.Encryption = buf.Bytes() + + _, err = p.sqlDb.Exec(`REPLACE INTO FileMetaData (Id, Name, Size, SHA1, ExpireAt, SizeBytes, + DownloadsRemaining, DownloadCount, PasswordHash, HotlinkId, ContentType, AwsBucket, Encryption, + UnlimitedDownloads, UnlimitedTime, UserId, UploadDate, PendingDeletion, UploadRequestId) + VALUES (?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?)`, + newData.Id, newData.Name, newData.Size, newData.SHA1, newData.ExpireAt, newData.SizeBytes, + newData.DownloadsRemaining, newData.DownloadCount, newData.PasswordHash, newData.HotlinkId, newData.ContentType, + newData.AwsBucket, newData.Encryption, newData.UnlimitedDownloads, newData.UnlimitedTime, newData.UserId, newData.UploadDate, + newData.PendingDeletion, newData.UploadRequestId) + helper.Check(err) +} + +// IncreaseDownloadCount increases the download count of a file atomically +func (p DatabaseProvider) IncreaseDownloadCount(id string, decreaseRemainingDownloads bool) { + if decreaseRemainingDownloads { + _, err := p.sqlDb.Exec(`UPDATE FileMetaData SET DownloadCount = DownloadCount + 1, + DownloadsRemaining = DownloadsRemaining - 1 WHERE Id = ?`, id) + helper.Check(err) + } else { + _, err := p.sqlDb.Exec(`UPDATE FileMetaData SET DownloadCount = DownloadCount + 1 WHERE Id = ?`, id) + helper.Check(err) + } +} + +// GetDownloadsRemaining returns the remaining downloads of a file that does not implement UnlimitedDownloads +func (p DatabaseProvider) GetDownloadsRemaining(id string) int { + var downloadsRemaining int + row := p.sqlDb.QueryRow("SELECT DownloadsRemaining FROM FileMetaData WHERE Id = ?", id) + err := row.Scan(&downloadsRemaining) + if err != nil { + if errors.Is(err, sql.ErrNoRows) { + return 0 + } + helper.Check(err) + return downloadsRemaining + } + return downloadsRemaining +} + +// DeleteMetaData deletes information about a file +func (p DatabaseProvider) DeleteMetaData(id string) { + _, err := p.sqlDb.Exec("DELETE FROM FileMetaData WHERE Id = ?", id) + helper.Check(err) +} diff --git a/internal/configuration/database/provider/mariadb/sessions.go b/internal/configuration/database/provider/mariadb/sessions.go new file mode 100644 index 00000000..e7446ad0 --- /dev/null +++ b/internal/configuration/database/provider/mariadb/sessions.go @@ -0,0 +1,75 @@ +package mariadb + +import ( + "database/sql" + "errors" + "time" + + "github.com/forceu/gokapi/internal/helper" + "github.com/forceu/gokapi/internal/models" +) + +type schemaSessions struct { + Id string + RenewAt int64 + ValidUntil int64 + UserId int +} + +// GetSession returns the session with the given ID or false if not a valid ID +func (p DatabaseProvider) GetSession(id string) (models.Session, bool) { + var rowResult schemaSessions + row := p.sqlDb.QueryRow("SELECT * FROM Sessions WHERE Id = ?", id) + err := row.Scan(&rowResult.Id, &rowResult.RenewAt, &rowResult.ValidUntil, &rowResult.UserId) + if err != nil { + if errors.Is(err, sql.ErrNoRows) { + return models.Session{}, false + } + helper.Check(err) + return models.Session{}, false + } + result := models.Session{ + RenewAt: rowResult.RenewAt, + ValidUntil: rowResult.ValidUntil, + UserId: rowResult.UserId, + } + return result, true +} + +// SaveSession stores the given session. After the expiry passed, it will be deleted automatically +func (p DatabaseProvider) SaveSession(id string, session models.Session) { + newData := schemaSessions{ + Id: id, + RenewAt: session.RenewAt, + ValidUntil: session.ValidUntil, + UserId: session.UserId, + } + + _, err := p.sqlDb.Exec("REPLACE INTO Sessions (Id, RenewAt, ValidUntil, UserId) VALUES (?, ?, ?, ?)", + newData.Id, newData.RenewAt, newData.ValidUntil, newData.UserId) + helper.Check(err) +} + +// DeleteSession deletes a session with the given ID +func (p DatabaseProvider) DeleteSession(id string) { + _, err := p.sqlDb.Exec("DELETE FROM Sessions WHERE Id = ?", id) + helper.Check(err) +} + +// DeleteAllSessions logs all users out +func (p DatabaseProvider) DeleteAllSessions() { + //goland:noinspection SqlWithoutWhere + _, err := p.sqlDb.Exec("DELETE FROM Sessions") + helper.Check(err) +} + +// DeleteAllSessionsByUser logs the specific users out +func (p DatabaseProvider) DeleteAllSessionsByUser(userId int) { + _, err := p.sqlDb.Exec("DELETE FROM Sessions WHERE UserId = ?", userId) + helper.Check(err) +} + +func (p DatabaseProvider) cleanExpiredSessions() { + _, err := p.sqlDb.Exec("DELETE FROM Sessions WHERE Sessions.ValidUntil < ?", time.Now().Unix()) + helper.Check(err) +} diff --git a/internal/configuration/database/provider/mariadb/statistics.go b/internal/configuration/database/provider/mariadb/statistics.go new file mode 100644 index 00000000..5a21777f --- /dev/null +++ b/internal/configuration/database/provider/mariadb/statistics.go @@ -0,0 +1,55 @@ +package mariadb + +import ( + "database/sql" + "errors" + + "github.com/forceu/gokapi/internal/helper" +) + +const statIdTraffic = 1 +const statIdTrafficSince = 2 + +// GetStatTraffic returns the total traffic from statistics +func (p DatabaseProvider) GetStatTraffic() uint64 { + var result uint64 + row := p.sqlDb.QueryRow("SELECT Value FROM Statistics WHERE Type = ?", statIdTraffic) + err := row.Scan(&result) + if err != nil { + if errors.Is(err, sql.ErrNoRows) { + return 0 + } + helper.Check(err) + return 0 + } + return result +} + +// SaveStatTraffic stores the total traffic +func (p DatabaseProvider) SaveStatTraffic(totalTraffic uint64) { + _, err := p.sqlDb.Exec(`INSERT INTO Statistics (Type, Value) VALUES (?, ?) + ON DUPLICATE KEY UPDATE Value = ?`, statIdTraffic, totalTraffic, totalTraffic) + helper.Check(err) +} + +// SaveTrafficSince stores the beginning of traffic counting +func (p DatabaseProvider) SaveTrafficSince(since int64) { + _, err := p.sqlDb.Exec(`INSERT INTO Statistics (Type, Value) VALUES (?, ?) + ON DUPLICATE KEY UPDATE Value = ?`, statIdTrafficSince, since, since) + helper.Check(err) +} + +// GetTrafficSince gets the beginning of traffic counting +func (p DatabaseProvider) GetTrafficSince() (int64, bool) { + var result int64 + row := p.sqlDb.QueryRow("SELECT Value FROM Statistics WHERE Type = ?", statIdTrafficSince) + err := row.Scan(&result) + if err != nil { + if errors.Is(err, sql.ErrNoRows) { + return 0, false + } + helper.Check(err) + return 0, false + } + return result, true +} diff --git a/internal/configuration/database/provider/mariadb/users.go b/internal/configuration/database/provider/mariadb/users.go new file mode 100644 index 00000000..64b50c3d --- /dev/null +++ b/internal/configuration/database/provider/mariadb/users.go @@ -0,0 +1,110 @@ +package mariadb + +import ( + "database/sql" + "errors" + "time" + + "github.com/forceu/gokapi/internal/helper" + "github.com/forceu/gokapi/internal/models" +) + +type schemaUser struct { + Id int + Name string + Password sql.NullString + Permissions models.UserPermission + UserLevel models.UserRank + LastOnline int64 + ResetPassword int +} + +func (s schemaUser) ToUser() models.User { + pw := "" + if s.Password.Valid { + pw = s.Password.String + } + return models.User{ + Id: s.Id, + Name: s.Name, + Permissions: s.Permissions, + UserLevel: s.UserLevel, + LastOnline: s.LastOnline, + Password: pw, + ResetPassword: s.ResetPassword == 1, + } +} + +// GetAllUsers returns a map with all users +func (p DatabaseProvider) GetAllUsers() []models.User { + var result []models.User + rows, err := p.sqlDb.Query("SELECT * FROM Users ORDER BY Userlevel, LastOnline DESC, Name") + helper.Check(err) + defer rows.Close() + for rows.Next() { + row := schemaUser{} + err = rows.Scan(&row.Id, &row.Name, &row.Password, &row.Permissions, &row.UserLevel, &row.LastOnline, &row.ResetPassword) + helper.Check(err) + result = append(result, row.ToUser()) + } + return result +} + +func (p DatabaseProvider) getUserWithConstraint(isName bool, searchValue any) (models.User, bool) { + rowResult := schemaUser{} + query := "SELECT * FROM Users WHERE Id = ?" + if isName { + query = "SELECT * FROM Users WHERE Name = ?" + } + row := p.sqlDb.QueryRow(query, searchValue) + err := row.Scan(&rowResult.Id, &rowResult.Name, &rowResult.Password, &rowResult.Permissions, &rowResult.UserLevel, &rowResult.LastOnline, &rowResult.ResetPassword) + if err != nil { + if errors.Is(err, sql.ErrNoRows) { + return models.User{}, false + } + helper.Check(err) + return models.User{}, false + } + user := rowResult.ToUser() + return user, true +} + +// GetUser returns a models.User if valid or false if the ID is not valid +func (p DatabaseProvider) GetUser(id int) (models.User, bool) { + return p.getUserWithConstraint(false, id) +} + +// GetUserByName returns a models.User if valid or false if the name is not valid +func (p DatabaseProvider) GetUserByName(username string) (models.User, bool) { + return p.getUserWithConstraint(true, username) +} + +// SaveUser saves a user to the database. If isNewUser is true, a new Id will be generated +func (p DatabaseProvider) SaveUser(user models.User, isNewUser bool) { + resetpw := 0 + if user.ResetPassword { + resetpw = 1 + } + if isNewUser { + _, err := p.sqlDb.Exec("INSERT INTO Users (Name, Password, Permissions, Userlevel, LastOnline, ResetPassword) VALUES (?, ?, ?, ?, ?, ?)", + user.Name, user.Password, user.Permissions, user.UserLevel, user.LastOnline, resetpw) + helper.Check(err) + } else { + _, err := p.sqlDb.Exec("REPLACE INTO Users (Id, Name, Password, Permissions, Userlevel, LastOnline, ResetPassword) VALUES (?, ?, ?, ?, ?, ?, ?)", + user.Id, user.Name, user.Password, user.Permissions, user.UserLevel, user.LastOnline, resetpw) + helper.Check(err) + } +} + +// UpdateUserLastOnline writes the last online time to the database +func (p DatabaseProvider) UpdateUserLastOnline(id int) { + timeNow := time.Now().Unix() + _, err := p.sqlDb.Exec("UPDATE Users SET LastOnline = ? WHERE Id = ?", timeNow, id) + helper.Check(err) +} + +// DeleteUser deletes a user with the given ID +func (p DatabaseProvider) DeleteUser(id int) { + _, err := p.sqlDb.Exec("DELETE FROM Users WHERE Id = ?", id) + helper.Check(err) +} diff --git a/internal/configuration/database/provider/postgres/Postgres.go b/internal/configuration/database/provider/postgres/Postgres.go new file mode 100644 index 00000000..424cea7d --- /dev/null +++ b/internal/configuration/database/provider/postgres/Postgres.go @@ -0,0 +1,212 @@ +package postgres + +import ( + "database/sql" + "errors" + "fmt" + + "github.com/forceu/gokapi/internal/helper" + "github.com/forceu/gokapi/internal/models" + // Required for the postgres driver + _ "github.com/jackc/pgx/v5/stdlib" +) + +// DatabaseProvider contains the database instance +type DatabaseProvider struct { + sqlDb *sql.DB +} + +// DatabaseSchemeVersion contains the version number to be expected from the current database. If lower, an upgrade will be performed +const DatabaseSchemeVersion = 1 + +// New returns an instance +func New(dbConfig models.DbConnection) (DatabaseProvider, error) { + return DatabaseProvider{}.init(dbConfig) +} + +// GetType returns 3, for being a PostgreSQL interface +func (p DatabaseProvider) GetType() int { + return 3 // dbabstraction.TypePostgres +} + +// Upgrade migrates the DB to a new Gokapi version, if required +func (p DatabaseProvider) Upgrade(currentDbVersion int) { + // No upgrade steps yet - this provider was introduced at schema version 1 +} + +// GetDbVersion gets the version number of the database +func (p DatabaseProvider) GetDbVersion() int { + var version int + row := p.sqlDb.QueryRow("SELECT version FROM schemaversion LIMIT 1") + err := row.Scan(&version) + if err != nil { + if errors.Is(err, sql.ErrNoRows) { + return 0 + } + helper.Check(err) + } + return version +} + +// SetDbVersion sets the version number of the database +func (p DatabaseProvider) SetDbVersion(newVersion int) { + _, err := p.sqlDb.Exec("DELETE FROM schemaversion") + helper.Check(err) + _, err = p.sqlDb.Exec("INSERT INTO schemaversion (version) VALUES ($1)", newVersion) + helper.Check(err) +} + +// GetSchemaVersion returns the version number, which the database should be at if fully upgraded +func (p DatabaseProvider) GetSchemaVersion() int { + return DatabaseSchemeVersion +} + +// Init connects to the database and creates the table structure, if necessary +func (p DatabaseProvider) init(dbConfig models.DbConnection) (DatabaseProvider, error) { + if dbConfig.HostUrl == "" { + return DatabaseProvider{}, errors.New("empty database url was provided") + } + if dbConfig.DatabaseName == "" { + return DatabaseProvider{}, errors.New("no database name was provided") + } + if p.sqlDb == nil { + dsn := fmt.Sprintf("postgres://%s:%s@%s/%s?sslmode=disable", dbConfig.Username, dbConfig.Password, dbConfig.HostUrl, dbConfig.DatabaseName) + var err error + p.sqlDb, err = sql.Open("pgx", dsn) + if err != nil { + return DatabaseProvider{}, err + } + p.sqlDb.SetMaxOpenConns(10) + p.sqlDb.SetMaxIdleConns(10) + + err = p.sqlDb.Ping() + if err != nil { + return DatabaseProvider{}, err + } + + exists, err := p.tableExists("filemetadata") + if err != nil { + return DatabaseProvider{}, err + } + if !exists { + return p, p.createNewDatabase() + } + return p, nil + } + return p, nil +} + +// Close the database connection +func (p DatabaseProvider) Close() { + if p.sqlDb != nil { + err := p.sqlDb.Close() + if err != nil { + fmt.Println(err) + } + } + p.sqlDb = nil +} + +// RunGarbageCollection runs the databases GC +func (p DatabaseProvider) RunGarbageCollection() { + p.cleanExpiredSessions() + p.cleanApiKeys() +} + +func (p DatabaseProvider) tableExists(tableName string) (bool, error) { + var count int + row := p.sqlDb.QueryRow("SELECT COUNT(*) FROM information_schema.tables WHERE table_schema = current_schema() AND table_name = $1", tableName) + err := row.Scan(&count) + if err != nil { + return false, err + } + return count > 0, nil +} + +func (p DatabaseProvider) createNewDatabase() error { + statements := []string{ + `CREATE TABLE apikeys ( + id VARCHAR(64) NOT NULL PRIMARY KEY, + friendlyname VARCHAR(255) NOT NULL, + lastused BIGINT NOT NULL, + permissions INT NOT NULL DEFAULT 0, + expiry BIGINT NOT NULL DEFAULT 0, + issystemkey SMALLINT NOT NULL DEFAULT 0, + userid INT NOT NULL, + publicid VARCHAR(64) NOT NULL UNIQUE, + uploadrequestid VARCHAR(64) NOT NULL + )`, + `CREATE TABLE e2econfig ( + id SERIAL PRIMARY KEY, + config BYTEA NOT NULL, + userid INT NOT NULL UNIQUE + )`, + `CREATE TABLE filemetadata ( + id VARCHAR(64) NOT NULL PRIMARY KEY, + name VARCHAR(1024) NOT NULL, + size VARCHAR(64) NOT NULL, + sha1 VARCHAR(64) NOT NULL, + expireat BIGINT NOT NULL, + sizebytes BIGINT NOT NULL, + downloadsremaining INT NOT NULL, + downloadcount INT NOT NULL, + passwordhash VARCHAR(255) NOT NULL, + hotlinkid VARCHAR(255) NOT NULL, + contenttype VARCHAR(255) NOT NULL, + awsbucket VARCHAR(255) NOT NULL, + encryption BYTEA NOT NULL, + unlimiteddownloads SMALLINT NOT NULL, + unlimitedtime SMALLINT NOT NULL, + userid INT NOT NULL, + uploaddate BIGINT NOT NULL, + pendingdeletion BIGINT NOT NULL, + uploadrequestid VARCHAR(64) NOT NULL + )`, + `CREATE TABLE hotlinks ( + id VARCHAR(255) NOT NULL PRIMARY KEY, + fileid VARCHAR(64) NOT NULL UNIQUE + )`, + `CREATE TABLE sessions ( + id VARCHAR(255) NOT NULL PRIMARY KEY, + renewat BIGINT NOT NULL, + validuntil BIGINT NOT NULL, + userid INT NOT NULL + )`, + `CREATE TABLE users ( + id SERIAL PRIMARY KEY, + name VARCHAR(255) NOT NULL UNIQUE, + password VARCHAR(255), + permissions INT NOT NULL, + userlevel INT NOT NULL, + lastonline BIGINT NOT NULL DEFAULT 0, + resetpassword SMALLINT NOT NULL DEFAULT 0 + )`, + `CREATE TABLE uploadrequests ( + id VARCHAR(64) NOT NULL PRIMARY KEY, + name VARCHAR(255) NOT NULL, + userid INT NOT NULL, + expiry BIGINT NOT NULL, + maxfiles INT NOT NULL, + maxsize INT NOT NULL, + creation BIGINT NOT NULL, + apikey VARCHAR(255) NOT NULL UNIQUE, + note TEXT NOT NULL + )`, + `CREATE TABLE statistics ( + id SERIAL PRIMARY KEY, + type INT NOT NULL UNIQUE, + value BIGINT + )`, + `CREATE TABLE schemaversion ( + version INT NOT NULL + )`, + } + for _, statement := range statements { + _, err := p.sqlDb.Exec(statement) + if err != nil { + return err + } + } + p.SetDbVersion(DatabaseSchemeVersion) + return nil +} diff --git a/internal/configuration/database/provider/postgres/Postgres_test.go b/internal/configuration/database/provider/postgres/Postgres_test.go new file mode 100644 index 00000000..1b1244dc --- /dev/null +++ b/internal/configuration/database/provider/postgres/Postgres_test.go @@ -0,0 +1,719 @@ +//go:build test && postgrestest + +// Package postgres tests require a real PostgreSQL server. They are gated behind the +// "postgrestest" build tag (not part of the default "test" tag set) and configured via +// environment variables, since there is no pure-Go in-process mock for PostgreSQL +// (unlike SQLite, which is just a file, or Redis, which uses miniredis): +// +// GOKAPI_POSTGRES_HOST=127.0.0.1:5432 +// GOKAPI_POSTGRES_DBNAME=gokapi_test +// GOKAPI_POSTGRES_USER=gokapi +// GOKAPI_POSTGRES_PASSWORD=secret +// +// Run with: go test ./internal/configuration/database/provider/postgres/... --tags=test,postgrestest +// +// The target database must already exist; the tables in it are dropped and recreated on +// every run, so use a dedicated, disposable database - never point this at production data. +package postgres + +import ( + "fmt" + "math" + "os" + "slices" + "sync" + "testing" + "time" + + "github.com/forceu/gokapi/internal/helper" + "github.com/forceu/gokapi/internal/models" + "github.com/forceu/gokapi/internal/test" +) + +var config = models.DbConnection{ + HostUrl: envOrDefault("GOKAPI_POSTGRES_HOST", "127.0.0.1:5432"), + DatabaseName: envOrDefault("GOKAPI_POSTGRES_DBNAME", "gokapi_test"), + Username: envOrDefault("GOKAPI_POSTGRES_USER", "gokapi"), + Password: envOrDefault("GOKAPI_POSTGRES_PASSWORD", "secret"), + Type: 3, // dbabstraction.TypePostgres +} + +func envOrDefault(key, fallback string) string { + if val, ok := os.LookupEnv(key); ok { + return val + } + return fallback +} + +var dbInstance DatabaseProvider + +func dropAllTables() error { + instance, err := New(config) + if err != nil { + return err + } + tables := []string{"apikeys", "e2econfig", "filemetadata", "hotlinks", "sessions", "users", "uploadrequests", "statistics", "schemaversion"} + for _, table := range tables { + _, err = instance.sqlDb.Exec("DROP TABLE IF EXISTS " + table) + if err != nil { + return err + } + } + instance.Close() + return nil +} + +func TestMain(m *testing.M) { + err := dropAllTables() + if err != nil { + fmt.Println("Could not reset PostgreSQL test database:", err) + os.Exit(1) + } + exitVal := m.Run() + os.Exit(exitVal) +} + +func TestInit(t *testing.T) { + instance, err := New(config) + test.IsNil(t, err) + instance.Close() + + _, err = New(models.DbConnection{HostUrl: "", DatabaseName: "gokapi_test", Type: 3}) + test.IsNotNil(t, err) + _, err = New(models.DbConnection{HostUrl: config.HostUrl, DatabaseName: "", Type: 3}) + test.IsNotNil(t, err) +} + +func TestClose(t *testing.T) { + instance, err := New(config) + test.IsNil(t, err) + instance.Close() + instance, err = New(config) + test.IsNil(t, err) + dbInstance = instance +} + +func TestDatabaseProvider_GetType(t *testing.T) { + test.IsEqualInt(t, dbInstance.GetType(), 3) +} + +func TestDatabaseProvider_GetDbVersion(t *testing.T) { + version := dbInstance.GetDbVersion() + test.IsEqualInt(t, version, DatabaseSchemeVersion) + dbInstance.SetDbVersion(99) + test.IsEqualInt(t, dbInstance.GetDbVersion(), 99) + dbInstance.SetDbVersion(DatabaseSchemeVersion) +} + +func TestDatabaseProvider_GetSchemaVersion(t *testing.T) { + test.IsEqualInt(t, dbInstance.GetSchemaVersion(), DatabaseSchemeVersion) +} + +func TestMetaData(t *testing.T) { + files := dbInstance.GetAllMetadata() + test.IsEqualInt(t, len(files), 0) + + dbInstance.SaveMetaData(models.File{Id: "testfile", Name: "test.txt", ExpireAt: time.Now().Add(time.Hour).Unix()}) + files = dbInstance.GetAllMetadata() + test.IsEqualInt(t, len(files), 1) + test.IsEqualString(t, files["testfile"].Name, "test.txt") + + file, ok := dbInstance.GetMetaDataById("testfile") + test.IsEqualBool(t, ok, true) + test.IsEqualString(t, file.Id, "testfile") + _, ok = dbInstance.GetMetaDataById("invalid") + test.IsEqualBool(t, ok, false) + + test.IsEqualInt(t, len(dbInstance.GetAllMetadata()), 1) + dbInstance.DeleteMetaData("invalid") + test.IsEqualInt(t, len(dbInstance.GetAllMetadata()), 1) + + test.IsEqualBool(t, file.UnlimitedDownloads, false) + test.IsEqualBool(t, file.UnlimitedTime, false) + + dbInstance.DeleteMetaData("testfile") + test.IsEqualInt(t, len(dbInstance.GetAllMetadata()), 0) + + dbInstance.SaveMetaData(models.File{ + Id: "test2", + Name: "test2", + UnlimitedDownloads: true, + UnlimitedTime: false, + }) + + file, ok = dbInstance.GetMetaDataById("test2") + test.IsEqualBool(t, ok, true) + test.IsEqualBool(t, file.UnlimitedDownloads, true) + test.IsEqualBool(t, file.UnlimitedTime, false) + + dbInstance.SaveMetaData(models.File{ + Id: "test3", + Name: "test3", + DownloadsRemaining: 4, + UnlimitedDownloads: false, + UnlimitedTime: true, + }) + file, ok = dbInstance.GetMetaDataById("test3") + test.IsEqualBool(t, ok, true) + test.IsEqualBool(t, file.UnlimitedDownloads, false) + test.IsEqualBool(t, file.UnlimitedTime, true) + test.IsEqualInt(t, file.DownloadsRemaining, 4) + remaining := dbInstance.GetDownloadsRemaining(file.Id) + test.IsEqualInt(t, remaining, 4) + + dbInstance.DeleteMetaData("test2") + dbInstance.DeleteMetaData("test3") +} + +func TestHotlink(t *testing.T) { + dbInstance.SaveHotlink(models.File{Id: "testfile", Name: "test.txt", HotlinkId: "testlink", ExpireAt: time.Now().Add(time.Hour).Unix()}) + + hotlink, ok := dbInstance.GetHotlink("testlink") + test.IsEqualBool(t, ok, true) + test.IsEqualString(t, hotlink, "testfile") + _, ok = dbInstance.GetHotlink("invalid") + test.IsEqualBool(t, ok, false) + + dbInstance.DeleteHotlink("invalid") + _, ok = dbInstance.GetHotlink("testlink") + test.IsEqualBool(t, ok, true) + dbInstance.DeleteHotlink("testlink") + _, ok = dbInstance.GetHotlink("testlink") + test.IsEqualBool(t, ok, false) + + dbInstance.SaveHotlink(models.File{Id: "testfile", Name: "test.txt", HotlinkId: "testlink", ExpireAt: 0, UnlimitedTime: true}) + hotlink, ok = dbInstance.GetHotlink("testlink") + test.IsEqualBool(t, ok, true) + test.IsEqualString(t, hotlink, "testfile") + + dbInstance.SaveHotlink(models.File{Id: "file2", Name: "file2.txt", HotlinkId: "link2", ExpireAt: time.Now().Add(time.Hour).Unix()}) + dbInstance.SaveHotlink(models.File{Id: "file3", Name: "file3.txt", HotlinkId: "link3", ExpireAt: time.Now().Add(time.Hour).Unix()}) + + hotlinks := dbInstance.GetAllHotlinks() + test.IsEqualInt(t, len(hotlinks), 3) + test.IsEqualBool(t, slices.Contains(hotlinks, "testlink"), true) + test.IsEqualBool(t, slices.Contains(hotlinks, "link2"), true) + test.IsEqualBool(t, slices.Contains(hotlinks, "link3"), true) + dbInstance.DeleteHotlink("") + hotlinks = dbInstance.GetAllHotlinks() + test.IsEqualInt(t, len(hotlinks), 3) + + dbInstance.DeleteHotlink("testlink") + dbInstance.DeleteHotlink("link2") + dbInstance.DeleteHotlink("link3") +} + +func TestDatabaseProvider_IncreaseDownloadCount(t *testing.T) { + newFile := models.File{ + Id: "newFileId", + Name: "newFileName", + Size: "3GB", + SHA1: "newSHA1", + PasswordHash: "newPassword", + HotlinkId: "newHotlink", + ContentType: "newContent", + AwsBucket: "newAws", + ExpireAt: 123456, + SizeBytes: 456789, + DownloadsRemaining: 11, + DownloadCount: 2, + Encryption: models.EncryptionInfo{ + IsEncrypted: true, + IsEndToEndEncrypted: true, + DecryptionKey: []byte("newDecryptionKey"), + Nonce: []byte("newDecryptionNonce"), + }, + UnlimitedDownloads: true, + UnlimitedTime: true, + } + dbInstance.SaveMetaData(newFile) + dbInstance.IncreaseDownloadCount(newFile.Id, false) + retrievedFile, ok := dbInstance.GetMetaDataById(newFile.Id) + test.IsEqualBool(t, ok, true) + test.IsEqualInt(t, retrievedFile.DownloadCount, 3) + test.IsEqualInt(t, retrievedFile.DownloadsRemaining, 11) + newFile.DownloadCount = 3 + test.IsEqual(t, retrievedFile, newFile) + + dbInstance.IncreaseDownloadCount(newFile.Id, true) + retrievedFile, ok = dbInstance.GetMetaDataById(newFile.Id) + test.IsEqualBool(t, ok, true) + test.IsEqualInt(t, retrievedFile.DownloadCount, 4) + test.IsEqualInt(t, retrievedFile.DownloadsRemaining, 10) + newFile.DownloadCount = 4 + newFile.DownloadsRemaining = 10 + test.IsEqual(t, retrievedFile, newFile) + dbInstance.DeleteMetaData(newFile.Id) +} + +func TestApiKey(t *testing.T) { + key1 := models.ApiKey{ + Id: "newkey", + FriendlyName: "New Key", + LastUsed: 100, + Permissions: 20, + PublicId: "_n3wkey", + Expiry: 0, + IsSystemKey: false, + UserId: 5, + } + key2 := models.ApiKey{ + Id: "newkey2", + FriendlyName: "New Key2", + PublicId: "_n3wkey2", + Expiry: 17362039396, + LastUsed: 200, + Permissions: 40, + IsSystemKey: true, + UserId: 10, + } + dbInstance.SaveApiKey(key1) + dbInstance.SaveApiKey(key2) + dbInstance.SaveApiKey(models.ApiKey{ + Id: "expiredKey", + PublicId: "expiredKey", + FriendlyName: "expiredKey", + Expiry: 1, + }) + + keys := dbInstance.GetAllApiKeys() + test.IsEqualInt(t, len(keys), 2) + test.IsEqual(t, keys["newkey"], key1) + test.IsEqual(t, keys["newkey2"], key2) + dbInstance.DeleteApiKey("newkey2") + test.IsEqualInt(t, len(dbInstance.GetAllApiKeys()), 1) + + key, ok := dbInstance.GetApiKey("newkey") + test.IsEqualBool(t, ok, true) + test.IsEqual(t, key, key1) + _, ok = dbInstance.GetApiKey("newkey2") + test.IsEqualBool(t, ok, false) + + dbInstance.SaveApiKey(models.ApiKey{ + Id: "newkey", + FriendlyName: "Old Key", + LastUsed: 100, + }) + key, ok = dbInstance.GetApiKey("newkey") + test.IsEqualBool(t, ok, true) + test.IsEqualString(t, key.FriendlyName, "Old Key") + + dbInstance.DeleteApiKey("newkey") + dbInstance.DeleteApiKey("expiredKey") +} + +func TestSession(t *testing.T) { + renewAt := time.Now().Add(1 * time.Hour).Unix() + dbInstance.SaveSession("newsession", models.Session{ + RenewAt: renewAt, + ValidUntil: time.Now().Add(2 * time.Hour).Unix(), + }) + + session, ok := dbInstance.GetSession("newsession") + test.IsEqualBool(t, ok, true) + test.IsEqualBool(t, session.RenewAt == renewAt, true) + + dbInstance.DeleteSession("newsession") + _, ok = dbInstance.GetSession("newsession") + test.IsEqualBool(t, ok, false) + + dbInstance.SaveSession("newsession", models.Session{ + RenewAt: renewAt, + ValidUntil: time.Now().Add(2 * time.Hour).Unix(), + }) + + dbInstance.SaveSession("anothersession", models.Session{ + RenewAt: renewAt, + ValidUntil: time.Now().Add(2 * time.Hour).Unix(), + }) + _, ok = dbInstance.GetSession("newsession") + test.IsEqualBool(t, ok, true) + _, ok = dbInstance.GetSession("anothersession") + test.IsEqualBool(t, ok, true) + + dbInstance.DeleteAllSessions() + _, ok = dbInstance.GetSession("newsession") + test.IsEqualBool(t, ok, false) + _, ok = dbInstance.GetSession("anothersession") + test.IsEqualBool(t, ok, false) + + session = models.Session{ + RenewAt: 2147483645, + ValidUntil: 2147483645, + UserId: 20, + } + dbInstance.SaveSession("sess_user1", session) + dbInstance.SaveSession("sess_user2", session) + dbInstance.SaveSession("sess_user3", session) + session.UserId = 40 + dbInstance.SaveSession("sess_user4", session) + _, ok = dbInstance.GetSession("sess_user1") + test.IsEqualBool(t, ok, true) + _, ok = dbInstance.GetSession("sess_user2") + test.IsEqualBool(t, ok, true) + _, ok = dbInstance.GetSession("sess_user3") + test.IsEqualBool(t, ok, true) + _, ok = dbInstance.GetSession("sess_user4") + test.IsEqualBool(t, ok, true) + dbInstance.DeleteAllSessionsByUser(20) + _, ok = dbInstance.GetSession("sess_user1") + test.IsEqualBool(t, ok, false) + _, ok = dbInstance.GetSession("sess_user2") + test.IsEqualBool(t, ok, false) + _, ok = dbInstance.GetSession("sess_user3") + test.IsEqualBool(t, ok, false) + _, ok = dbInstance.GetSession("sess_user4") + test.IsEqualBool(t, ok, true) + dbInstance.DeleteAllSessions() +} + +func TestFileRequest(t *testing.T) { + req1 := models.FileRequest{ + Id: "req1", + Name: "New file request", + UserId: 45564, + ApiKey: "123", + CreationDate: time.Now().Unix(), + } + dbInstance.SaveFileRequest(req1) + + request, ok := dbInstance.GetFileRequest("req1") + test.IsEqualBool(t, ok, true) + test.IsEqualString(t, request.Id, "req1") + test.IsEqualString(t, request.Name, "New file request") + + _, ok = dbInstance.GetFileRequest("invalid") + test.IsEqualBool(t, ok, false) + + _, ok = dbInstance.GetFileRequest("") + test.IsEqualBool(t, ok, false) + + dbInstance.DeleteFileRequest(models.FileRequest{Id: "invalid"}) + _, ok = dbInstance.GetFileRequest("req1") + test.IsEqualBool(t, ok, true) + + dbInstance.DeleteFileRequest(req1) + _, ok = dbInstance.GetFileRequest("req1") + test.IsEqualBool(t, ok, false) + + req2 := models.FileRequest{ + Id: "req2", + UserId: 45564, + Name: "file2.txt", + ApiKey: "456", + CreationDate: time.Now().Add(-time.Minute).Unix(), + } + req3 := models.FileRequest{ + Id: "req3", + Name: "file3.txt", + UserId: 45564, + ApiKey: "789", + CreationDate: time.Now().Add(-2 * time.Minute).Unix(), + } + + dbInstance.SaveFileRequest(req1) + dbInstance.SaveFileRequest(req2) + dbInstance.SaveFileRequest(req3) + + requests := dbInstance.GetAllFileRequests() + test.IsEqualInt(t, len(requests), 3) + + ids := []string{requests[0].Id, requests[1].Id, requests[2].Id} + test.IsEqualBool(t, slices.Contains(ids, "req1"), true) + test.IsEqualBool(t, slices.Contains(ids, "req2"), true) + test.IsEqualBool(t, slices.Contains(ids, "req3"), true) + + test.IsEqualBool(t, requests[0].CreationDate >= requests[1].CreationDate, true) + test.IsEqualBool(t, requests[1].CreationDate >= requests[2].CreationDate, true) + + dbInstance.DeleteFileRequest(req1) + dbInstance.DeleteFileRequest(req2) + dbInstance.DeleteFileRequest(req3) +} + +func TestGarbageCollectionSessions(t *testing.T) { + dbInstance.SaveSession("todelete1", models.Session{ + RenewAt: time.Now().Add(-10 * time.Second).Unix(), + ValidUntil: time.Now().Add(-10 * time.Second).Unix(), + }) + dbInstance.SaveSession("todelete2", models.Session{ + RenewAt: time.Now().Add(10 * time.Second).Unix(), + ValidUntil: time.Now().Add(-10 * time.Second).Unix(), + }) + dbInstance.SaveSession("tokeep1", models.Session{ + RenewAt: time.Now().Add(-10 * time.Second).Unix(), + ValidUntil: time.Now().Add(10 * time.Second).Unix(), + }) + dbInstance.SaveSession("tokeep2", models.Session{ + RenewAt: time.Now().Add(10 * time.Second).Unix(), + ValidUntil: time.Now().Add(10 * time.Second).Unix(), + }) + for _, item := range []string{"todelete1", "todelete2", "tokeep1", "tokeep2"} { + _, result := dbInstance.GetSession(item) + test.IsEqualBool(t, result, true) + } + dbInstance.RunGarbageCollection() + for _, item := range []string{"todelete1", "todelete2"} { + _, result := dbInstance.GetSession(item) + test.IsEqualBool(t, result, false) + } + for _, item := range []string{"tokeep1", "tokeep2"} { + _, result := dbInstance.GetSession(item) + test.IsEqualBool(t, result, true) + } + dbInstance.DeleteAllSessions() +} + +func TestEnd2EndInfo(t *testing.T) { + info := dbInstance.GetEnd2EndInfo(4) + test.IsEqualInt(t, info.Version, 0) + test.IsEqualBool(t, info.HasBeenSetUp(), false) + + dbInstance.SaveEnd2EndInfo(models.E2EInfoEncrypted{ + Version: 1, + Nonce: []byte("testNonce1"), + Content: []byte("testContent1"), + AvailableFiles: nil, + }, 4) + + info = dbInstance.GetEnd2EndInfo(4) + test.IsEqualInt(t, info.Version, 1) + test.IsEqualBool(t, info.HasBeenSetUp(), true) + test.IsEqualByteSlice(t, info.Nonce, []byte("testNonce1")) + test.IsEqualByteSlice(t, info.Content, []byte("testContent1")) + test.IsEqualBool(t, len(info.AvailableFiles) == 0, true) + + dbInstance.SaveEnd2EndInfo(models.E2EInfoEncrypted{ + Version: 2, + Nonce: []byte("testNonce2"), + Content: []byte("testContent2"), + AvailableFiles: nil, + }, 4) + + info = dbInstance.GetEnd2EndInfo(4) + test.IsEqualInt(t, info.Version, 2) + test.IsEqualBool(t, info.HasBeenSetUp(), true) + test.IsEqualByteSlice(t, info.Nonce, []byte("testNonce2")) + test.IsEqualByteSlice(t, info.Content, []byte("testContent2")) + test.IsEqualBool(t, len(info.AvailableFiles) == 0, true) + + dbInstance.DeleteEnd2EndInfo(4) + info = dbInstance.GetEnd2EndInfo(4) + test.IsEqualInt(t, info.Version, 0) + test.IsEqualBool(t, info.HasBeenSetUp(), false) +} + +func TestUpdateTimeApiKey(t *testing.T) { + retrievedKey, ok := dbInstance.GetApiKey("key1") + test.IsEqualBool(t, ok, false) + test.IsEqualString(t, retrievedKey.Id, "") + + key := models.ApiKey{ + Id: "key1", + FriendlyName: "key1", + PublicId: "key1", + LastUsed: 100, + } + dbInstance.SaveApiKey(key) + key = models.ApiKey{ + Id: "key2", + FriendlyName: "key2", + PublicId: "key2", + LastUsed: 200, + } + dbInstance.SaveApiKey(key) + + retrievedKey, ok = dbInstance.GetApiKey("key1") + test.IsEqualBool(t, ok, true) + test.IsEqualString(t, retrievedKey.Id, "key1") + test.IsEqualInt64(t, retrievedKey.LastUsed, 100) + retrievedKey, ok = dbInstance.GetApiKey("key2") + test.IsEqualBool(t, ok, true) + test.IsEqualString(t, retrievedKey.Id, "key2") + test.IsEqualInt64(t, retrievedKey.LastUsed, 200) + + key.LastUsed = 300 + dbInstance.UpdateTimeApiKey(key) + + retrievedKey, ok = dbInstance.GetApiKey("key1") + test.IsEqualBool(t, ok, true) + test.IsEqualString(t, retrievedKey.Id, "key1") + test.IsEqualInt64(t, retrievedKey.LastUsed, 100) + retrievedKey, ok = dbInstance.GetApiKey("key2") + test.IsEqualBool(t, ok, true) + test.IsEqualString(t, retrievedKey.Id, "key2") + test.IsEqualInt64(t, retrievedKey.LastUsed, 300) + + dbInstance.SaveApiKey(models.ApiKey{ + Id: "publicTest", + PublicId: "publicId", + }) + _, ok = dbInstance.GetApiKey("publicTest") + test.IsEqualBool(t, ok, true) + _, ok = dbInstance.GetApiKey("publicId") + test.IsEqualBool(t, ok, false) + _, ok = dbInstance.GetApiKeyByPublicKey("publicTest") + test.IsEqualBool(t, ok, false) + keyName, ok := dbInstance.GetApiKeyByPublicKey("publicId") + test.IsEqualBool(t, ok, true) + test.IsEqualString(t, keyName, "publicTest") + + dbInstance.DeleteApiKey("key1") + dbInstance.DeleteApiKey("key2") + dbInstance.DeleteApiKey("publicTest") +} + +func TestParallelConnectionsWritingAndReading(t *testing.T) { + var wg sync.WaitGroup + + simulatedConnection := func(t *testing.T) { + file := models.File{ + Id: helper.GenerateRandomString(10), + Name: helper.GenerateRandomString(10), + Size: "10B", + SHA1: "1289423794287598237489", + ExpireAt: math.MaxInt32, + SizeBytes: 10, + DownloadsRemaining: 10, + DownloadCount: 10, + Encryption: models.EncryptionInfo{}, + UnlimitedDownloads: false, + UnlimitedTime: false, + } + dbInstance.SaveMetaData(file) + retrievedFile, ok := dbInstance.GetMetaDataById(file.Id) + test.IsEqualBool(t, ok, true) + test.IsEqualString(t, retrievedFile.Name, file.Name) + dbInstance.DeleteMetaData(file.Id) + _, ok = dbInstance.GetMetaDataById(file.Id) + test.IsEqualBool(t, ok, false) + } + + for i := 1; i <= 100; i++ { + wg.Add(1) + go func() { + defer wg.Done() + simulatedConnection(t) + }() + } + wg.Wait() +} + +func TestParallelConnectionsReading(t *testing.T) { + var wg sync.WaitGroup + + dbInstance.SaveApiKey(models.ApiKey{ + Id: "readtest", + FriendlyName: "readtest", + LastUsed: 40000, + }) + simulatedConnection := func(t *testing.T) { + _, ok := dbInstance.GetApiKey("readtest") + test.IsEqualBool(t, ok, true) + } + + for i := 1; i <= 1000; i++ { + wg.Add(1) + go func() { + defer wg.Done() + simulatedConnection(t) + }() + } + wg.Wait() + dbInstance.DeleteApiKey("readtest") +} + +func TestStatistics(t *testing.T) { + test.IsEqualInt64(t, int64(dbInstance.GetStatTraffic()), 0) + dbInstance.SaveStatTraffic(1024) + test.IsEqualInt64(t, int64(dbInstance.GetStatTraffic()), 1024) + dbInstance.SaveStatTraffic(2048) + test.IsEqualInt64(t, int64(dbInstance.GetStatTraffic()), 2048) + + _, ok := dbInstance.GetTrafficSince() + test.IsEqualBool(t, ok, false) + dbInstance.SaveTrafficSince(12345) + since, ok := dbInstance.GetTrafficSince() + test.IsEqualBool(t, ok, true) + test.IsEqualInt64(t, since, 12345) + dbInstance.SaveTrafficSince(54321) + since, ok = dbInstance.GetTrafficSince() + test.IsEqualBool(t, ok, true) + test.IsEqualInt64(t, since, 54321) +} + +func TestUsers(t *testing.T) { + users := dbInstance.GetAllUsers() + test.IsEqualInt(t, len(users), 0) + user := models.User{ + Id: 2, + Name: "test", + Permissions: models.UserPermissionAll, + UserLevel: models.UserLevelUser, + LastOnline: 1337, + Password: "123456", + ResetPassword: true, + } + dbInstance.SaveUser(user, false) + retrievedUser, ok := dbInstance.GetUser(2) + test.IsEqualBool(t, ok, true) + test.IsEqual(t, retrievedUser, user) + users = dbInstance.GetAllUsers() + test.IsEqualInt(t, len(users), 1) + test.IsEqualInt(t, retrievedUser.Id, 2) + + _, ok = dbInstance.GetUser(0) + test.IsEqualBool(t, ok, false) + _, ok = dbInstance.GetUserByName("invalid") + test.IsEqualBool(t, ok, false) + retrievedUser, ok = dbInstance.GetUserByName("test") + test.IsEqualBool(t, ok, true) + test.IsEqual(t, retrievedUser, user) + + dbInstance.DeleteUser(2) + _, ok = dbInstance.GetUser(2) + test.IsEqualBool(t, ok, false) + + user = models.User{ + Id: 1000, + Name: "test2", + Permissions: models.UserPermissionNone, + UserLevel: models.UserLevelAdmin, + LastOnline: 1338, + Password: "1234568", + ResetPassword: true, + } + dbInstance.SaveUser(user, true) + _, ok = dbInstance.GetUser(1000) + test.IsEqualBool(t, ok, false) + retrievedUser, ok = dbInstance.GetUserByName("test2") + test.IsEqualBool(t, ok, true) + test.IsEqualBool(t, retrievedUser.Id == 1000, false) + user.Id = retrievedUser.Id + test.IsEqual(t, retrievedUser, user) + + dbInstance.UpdateUserLastOnline(retrievedUser.Id) + retrievedUser, ok = dbInstance.GetUser(retrievedUser.Id) + test.IsEqualBool(t, ok, true) + test.IsEqualBool(t, time.Now().Unix()-retrievedUser.LastOnline < 5, true) + test.IsEqualBool(t, time.Now().Unix()-retrievedUser.LastOnline > -1, true) + + user.Name = "test1" + dbInstance.SaveUser(user, true) + user.Name = "test3" + dbInstance.SaveUser(user, true) + user.Name = "test99" + user.UserLevel = models.UserLevelSuperAdmin + dbInstance.SaveUser(user, true) + user.Name = "test0" + user.UserLevel = models.UserLevelUser + dbInstance.SaveUser(user, true) + + users = dbInstance.GetAllUsers() + test.IsEqualInt(t, len(users), 5) + test.IsEqualString(t, users[0].Name, "test99") + test.IsEqualString(t, users[1].Name, "test2") + test.IsEqualString(t, users[2].Name, "test1") + test.IsEqualString(t, users[3].Name, "test3") + test.IsEqualString(t, users[4].Name, "test0") +} diff --git a/internal/configuration/database/provider/postgres/apikeys.go b/internal/configuration/database/provider/postgres/apikeys.go new file mode 100644 index 00000000..87791459 --- /dev/null +++ b/internal/configuration/database/provider/postgres/apikeys.go @@ -0,0 +1,130 @@ +package postgres + +import ( + "database/sql" + "errors" + "time" + + "github.com/forceu/gokapi/internal/helper" + "github.com/forceu/gokapi/internal/models" +) + +type schemaApiKeys struct { + Id string + FriendlyName string + LastUsed int64 + Permissions int + Expiry int64 + IsSystemKey int + UserId int + PublicId string + UploadRequestId string +} + +// currentTime is used in order to modify the current time for testing purposes in unit tests +var currentTime = func() time.Time { + return time.Now() +} + +// GetAllApiKeys returns a map with all API keys +func (p DatabaseProvider) GetAllApiKeys() map[string]models.ApiKey { + result := make(map[string]models.ApiKey) + + rows, err := p.sqlDb.Query("SELECT * FROM apikeys WHERE apikeys.expiry = 0 OR apikeys.expiry > $1", currentTime().Unix()) + helper.Check(err) + defer rows.Close() + for rows.Next() { + rowData := schemaApiKeys{} + err = rows.Scan(&rowData.Id, &rowData.FriendlyName, &rowData.LastUsed, &rowData.Permissions, &rowData.Expiry, + &rowData.IsSystemKey, &rowData.UserId, &rowData.PublicId, &rowData.UploadRequestId) + helper.Check(err) + result[rowData.Id] = models.ApiKey{ + Id: rowData.Id, + PublicId: rowData.PublicId, + FriendlyName: rowData.FriendlyName, + LastUsed: rowData.LastUsed, + Permissions: models.ApiPermission(rowData.Permissions), + Expiry: rowData.Expiry, + IsSystemKey: rowData.IsSystemKey == 1, + UserId: rowData.UserId, + UploadRequestId: rowData.UploadRequestId, + } + } + return result +} + +// GetApiKey returns a models.ApiKey if valid or false if the ID is not valid +func (p DatabaseProvider) GetApiKey(id string) (models.ApiKey, bool) { + var rowResult schemaApiKeys + row := p.sqlDb.QueryRow("SELECT * FROM apikeys WHERE id = $1", id) + err := row.Scan(&rowResult.Id, &rowResult.FriendlyName, &rowResult.LastUsed, &rowResult.Permissions, &rowResult.Expiry, + &rowResult.IsSystemKey, &rowResult.UserId, &rowResult.PublicId, &rowResult.UploadRequestId) + if err != nil { + if errors.Is(err, sql.ErrNoRows) { + return models.ApiKey{}, false + } + helper.Check(err) + return models.ApiKey{}, false + } + + result := models.ApiKey{ + Id: rowResult.Id, + PublicId: rowResult.PublicId, + FriendlyName: rowResult.FriendlyName, + LastUsed: rowResult.LastUsed, + Permissions: models.ApiPermission(rowResult.Permissions), + Expiry: rowResult.Expiry, + IsSystemKey: rowResult.IsSystemKey == 1, + UserId: rowResult.UserId, + UploadRequestId: rowResult.UploadRequestId, + } + + return result, true +} + +// GetApiKeyByPublicKey returns an API key by using the public key +func (p DatabaseProvider) GetApiKeyByPublicKey(publicKey string) (string, bool) { + var rowResult schemaApiKeys + row := p.sqlDb.QueryRow("SELECT id FROM apikeys WHERE publicid = $1 LIMIT 1", publicKey) + err := row.Scan(&rowResult.Id) + if err != nil { + if errors.Is(err, sql.ErrNoRows) { + return "", false + } + helper.Check(err) + return "", false + } + return rowResult.Id, true +} + +// SaveApiKey saves the API key to the database +func (p DatabaseProvider) SaveApiKey(apikey models.ApiKey) { + isSystemKey := 0 + if apikey.IsSystemKey { + isSystemKey = 1 + } + _, err := p.sqlDb.Exec(`INSERT INTO apikeys (id, friendlyname, lastused, permissions, expiry, issystemkey, userid, publicid, uploadrequestid) + VALUES ($1, $2, $3, $4, $5, $6, $7, $8, $9) + ON CONFLICT (id) DO UPDATE SET friendlyname = $2, lastused = $3, permissions = $4, expiry = $5, + issystemkey = $6, userid = $7, publicid = $8, uploadrequestid = $9`, + apikey.Id, apikey.FriendlyName, apikey.LastUsed, apikey.Permissions, apikey.Expiry, isSystemKey, apikey.UserId, apikey.PublicId, apikey.UploadRequestId) + helper.Check(err) +} + +// UpdateTimeApiKey writes the content of LastUsage to the database +func (p DatabaseProvider) UpdateTimeApiKey(apikey models.ApiKey) { + _, err := p.sqlDb.Exec("UPDATE apikeys SET lastused = $1 WHERE id = $2", + apikey.LastUsed, apikey.Id) + helper.Check(err) +} + +// DeleteApiKey deletes an API key with the given ID +func (p DatabaseProvider) DeleteApiKey(id string) { + _, err := p.sqlDb.Exec("DELETE FROM apikeys WHERE id = $1", id) + helper.Check(err) +} + +func (p DatabaseProvider) cleanApiKeys() { + _, err := p.sqlDb.Exec("DELETE FROM apikeys WHERE apikeys.expiry > 0 AND apikeys.expiry < $1", currentTime().Unix()) + helper.Check(err) +} diff --git a/internal/configuration/database/provider/postgres/e2econfig.go b/internal/configuration/database/provider/postgres/e2econfig.go new file mode 100644 index 00000000..81ad4125 --- /dev/null +++ b/internal/configuration/database/provider/postgres/e2econfig.go @@ -0,0 +1,58 @@ +package postgres + +import ( + "bytes" + "database/sql" + "encoding/gob" + "errors" + + "github.com/forceu/gokapi/internal/helper" + "github.com/forceu/gokapi/internal/models" +) + +type schemaE2EConfig struct { + Id int64 + Config []byte + UserId int +} + +// SaveEnd2EndInfo stores the encrypted e2e info +func (p DatabaseProvider) SaveEnd2EndInfo(info models.E2EInfoEncrypted, userId int) { + var buf bytes.Buffer + enc := gob.NewEncoder(&buf) + err := enc.Encode(info) + helper.Check(err) + + _, err = p.sqlDb.Exec(`INSERT INTO e2econfig (config, userid) VALUES ($1, $2) + ON CONFLICT (userid) DO UPDATE SET config = $1`, + buf.Bytes(), userId) + helper.Check(err) +} + +// GetEnd2EndInfo retrieves the encrypted e2e info +func (p DatabaseProvider) GetEnd2EndInfo(userId int) models.E2EInfoEncrypted { + result := models.E2EInfoEncrypted{} + rowResult := schemaE2EConfig{} + + row := p.sqlDb.QueryRow("SELECT config FROM e2econfig WHERE userid = $1", userId) + err := row.Scan(&rowResult.Config) + if err != nil { + if errors.Is(err, sql.ErrNoRows) { + return result + } + helper.Check(err) + return result + } + + buf := bytes.NewBuffer(rowResult.Config) + dec := gob.NewDecoder(buf) + err = dec.Decode(&result) + helper.Check(err) + return result +} + +// DeleteEnd2EndInfo resets the encrypted e2e info +func (p DatabaseProvider) DeleteEnd2EndInfo(userId int) { + _, err := p.sqlDb.Exec("DELETE FROM e2econfig WHERE userid = $1", userId) + helper.Check(err) +} diff --git a/internal/configuration/database/provider/postgres/filerequests.go b/internal/configuration/database/provider/postgres/filerequests.go new file mode 100644 index 00000000..80cf2919 --- /dev/null +++ b/internal/configuration/database/provider/postgres/filerequests.go @@ -0,0 +1,108 @@ +package postgres + +import ( + "database/sql" + "errors" + + "github.com/forceu/gokapi/internal/helper" + "github.com/forceu/gokapi/internal/models" +) + +type schemaFileRequests struct { + Id string + Name string + UserId int + Expiry int64 + MaxFiles int + MaxSize int + Creation int64 + ApiKey string + Note string +} + +// GetFileRequest returns the FileRequest or false if not found +func (p DatabaseProvider) GetFileRequest(id string) (models.FileRequest, bool) { + if id == "" { + return models.FileRequest{}, false + } + var rowResult schemaFileRequests + row := p.sqlDb.QueryRow("SELECT * FROM uploadrequests WHERE id = $1", id) + err := row.Scan(&rowResult.Id, &rowResult.Name, &rowResult.UserId, &rowResult.Expiry, + &rowResult.MaxFiles, &rowResult.MaxSize, &rowResult.Creation, &rowResult.ApiKey, &rowResult.Note) + if err != nil { + if errors.Is(err, sql.ErrNoRows) { + return models.FileRequest{}, false + } + helper.Check(err) + return models.FileRequest{}, false + } + result := models.FileRequest{ + Id: rowResult.Id, + Name: rowResult.Name, + UserId: rowResult.UserId, + MaxFiles: rowResult.MaxFiles, + MaxSize: rowResult.MaxSize, + Expiry: rowResult.Expiry, + CreationDate: rowResult.Creation, + ApiKey: rowResult.ApiKey, + Notes: rowResult.Note, + } + return result, true +} + +// GetAllFileRequests returns an array with all file requests, ordered by creation date +func (p DatabaseProvider) GetAllFileRequests() []models.FileRequest { + result := make([]models.FileRequest, 0) + rows, err := p.sqlDb.Query("SELECT * FROM uploadrequests ORDER BY creation DESC, name") + helper.Check(err) + defer rows.Close() + for rows.Next() { + rowData := schemaFileRequests{} + err = rows.Scan(&rowData.Id, &rowData.Name, &rowData.UserId, &rowData.Expiry, &rowData.MaxFiles, + &rowData.MaxSize, &rowData.Creation, &rowData.ApiKey, &rowData.Note) + helper.Check(err) + result = append(result, models.FileRequest{ + Id: rowData.Id, + Name: rowData.Name, + UserId: rowData.UserId, + MaxFiles: rowData.MaxFiles, + MaxSize: rowData.MaxSize, + Expiry: rowData.Expiry, + CreationDate: rowData.Creation, + ApiKey: rowData.ApiKey, + Notes: rowData.Note, + }) + } + return result +} + +// SaveFileRequest stores the file request associated with the file in the database +func (p DatabaseProvider) SaveFileRequest(request models.FileRequest) { + newData := schemaFileRequests{ + Id: request.Id, + Name: request.Name, + UserId: request.UserId, + MaxFiles: request.MaxFiles, + MaxSize: request.MaxSize, + Expiry: request.Expiry, + Creation: request.CreationDate, + ApiKey: request.ApiKey, + Note: request.Notes, + } + + _, err := p.sqlDb.Exec(`INSERT INTO uploadrequests (id, name, userid, expiry, maxfiles, maxsize, creation, apikey, note) + VALUES ($1, $2, $3, $4, $5, $6, $7, $8, $9) + ON CONFLICT (id) DO UPDATE SET name = $2, userid = $3, expiry = $4, maxfiles = $5, + maxsize = $6, creation = $7, apikey = $8, note = $9`, + newData.Id, newData.Name, newData.UserId, newData.Expiry, newData.MaxFiles, newData.MaxSize, newData.Creation, newData.ApiKey, newData.Note) + helper.Check(err) +} + +// DeleteFileRequest deletes a file request with the given ID +func (p DatabaseProvider) DeleteFileRequest(request models.FileRequest) { + if request.Id == "" { + return + } + _, err := p.sqlDb.Exec("DELETE FROM uploadrequests WHERE id = $1", request.Id) + helper.Check(err) +} diff --git a/internal/configuration/database/provider/postgres/hotlinks.go b/internal/configuration/database/provider/postgres/hotlinks.go new file mode 100644 index 00000000..e4951a24 --- /dev/null +++ b/internal/configuration/database/provider/postgres/hotlinks.go @@ -0,0 +1,66 @@ +package postgres + +import ( + "database/sql" + "errors" + + "github.com/forceu/gokapi/internal/helper" + "github.com/forceu/gokapi/internal/models" +) + +type schemaHotlinks struct { + Id string + FileId string +} + +// GetHotlink returns the id of the file associated or false if not found +func (p DatabaseProvider) GetHotlink(id string) (string, bool) { + var rowResult schemaHotlinks + row := p.sqlDb.QueryRow("SELECT fileid FROM hotlinks WHERE id = $1", id) + err := row.Scan(&rowResult.FileId) + if err != nil { + if errors.Is(err, sql.ErrNoRows) { + return "", false + } + helper.Check(err) + return "", false + } + return rowResult.FileId, true +} + +// GetAllHotlinks returns an array with all hotlink ids +func (p DatabaseProvider) GetAllHotlinks() []string { + ids := make([]string, 0) + rows, err := p.sqlDb.Query("SELECT id FROM hotlinks") + helper.Check(err) + defer rows.Close() + for rows.Next() { + rowData := schemaHotlinks{} + err = rows.Scan(&rowData.Id) + helper.Check(err) + ids = append(ids, rowData.Id) + } + return ids +} + +// SaveHotlink stores the hotlink associated with the file in the database +func (p DatabaseProvider) SaveHotlink(file models.File) { + newData := schemaHotlinks{ + Id: file.HotlinkId, + FileId: file.Id, + } + + _, err := p.sqlDb.Exec(`INSERT INTO hotlinks (id, fileid) VALUES ($1, $2) + ON CONFLICT (id) DO UPDATE SET fileid = $2`, + newData.Id, newData.FileId) + helper.Check(err) +} + +// DeleteHotlink deletes a hotlink with the given hotlink ID +func (p DatabaseProvider) DeleteHotlink(id string) { + if id == "" { + return + } + _, err := p.sqlDb.Exec("DELETE FROM hotlinks WHERE id = $1", id) + helper.Check(err) +} diff --git a/internal/configuration/database/provider/postgres/metadata.go b/internal/configuration/database/provider/postgres/metadata.go new file mode 100644 index 00000000..773a40af --- /dev/null +++ b/internal/configuration/database/provider/postgres/metadata.go @@ -0,0 +1,188 @@ +package postgres + +import ( + "bytes" + "database/sql" + "encoding/gob" + "errors" + + "github.com/forceu/gokapi/internal/helper" + "github.com/forceu/gokapi/internal/models" +) + +type schemaMetaData struct { + Id string + Name string + Size string + SHA1 string + ExpireAt int64 + SizeBytes int64 + DownloadsRemaining int + DownloadCount int + PasswordHash string + HotlinkId string + ContentType string + AwsBucket string + Encryption []byte + UnlimitedDownloads int + UnlimitedTime int + UserId int + UploadDate int64 + PendingDeletion int64 + UploadRequestId string +} + +func (rowData schemaMetaData) ToFileModel() (models.File, error) { + result := models.File{ + Id: rowData.Id, + Name: rowData.Name, + Size: rowData.Size, + SHA1: rowData.SHA1, + ExpireAt: rowData.ExpireAt, + SizeBytes: rowData.SizeBytes, + DownloadsRemaining: rowData.DownloadsRemaining, + DownloadCount: rowData.DownloadCount, + PasswordHash: rowData.PasswordHash, + HotlinkId: rowData.HotlinkId, + ContentType: rowData.ContentType, + AwsBucket: rowData.AwsBucket, + Encryption: models.EncryptionInfo{}, + UnlimitedDownloads: rowData.UnlimitedDownloads == 1, + UnlimitedTime: rowData.UnlimitedTime == 1, + UserId: rowData.UserId, + UploadDate: rowData.UploadDate, + PendingDeletion: rowData.PendingDeletion, + UploadRequestId: rowData.UploadRequestId, + } + + buf := bytes.NewBuffer(rowData.Encryption) + dec := gob.NewDecoder(buf) + err := dec.Decode(&result.Encryption) + return result, err +} + +// GetAllMetadata returns a map of all available files +func (p DatabaseProvider) GetAllMetadata() map[string]models.File { + result := make(map[string]models.File) + rows, err := p.sqlDb.Query("SELECT * FROM filemetadata") + helper.Check(err) + defer rows.Close() + for rows.Next() { + rowData := schemaMetaData{} + err = rows.Scan(&rowData.Id, &rowData.Name, &rowData.Size, &rowData.SHA1, &rowData.ExpireAt, &rowData.SizeBytes, + &rowData.DownloadsRemaining, &rowData.DownloadCount, &rowData.PasswordHash, &rowData.HotlinkId, &rowData.ContentType, + &rowData.AwsBucket, &rowData.Encryption, &rowData.UnlimitedDownloads, &rowData.UnlimitedTime, &rowData.UserId, + &rowData.UploadDate, &rowData.PendingDeletion, &rowData.UploadRequestId) + helper.Check(err) + var metaData models.File + metaData, err = rowData.ToFileModel() + helper.Check(err) + result[metaData.Id] = metaData + } + return result +} + +// GetMetaDataById returns a models.File from the ID passed or false if the id is not valid +func (p DatabaseProvider) GetMetaDataById(id string) (models.File, bool) { + result := models.File{} + rowData := schemaMetaData{} + + row := p.sqlDb.QueryRow("SELECT * FROM filemetadata WHERE id = $1", id) + err := row.Scan(&rowData.Id, &rowData.Name, &rowData.Size, &rowData.SHA1, &rowData.ExpireAt, &rowData.SizeBytes, + &rowData.DownloadsRemaining, &rowData.DownloadCount, &rowData.PasswordHash, + &rowData.HotlinkId, &rowData.ContentType, &rowData.AwsBucket, &rowData.Encryption, + &rowData.UnlimitedDownloads, &rowData.UnlimitedTime, &rowData.UserId, &rowData.UploadDate, + &rowData.PendingDeletion, &rowData.UploadRequestId) + if err != nil { + if errors.Is(err, sql.ErrNoRows) { + return result, false + } + helper.Check(err) + return result, false + } + result, err = rowData.ToFileModel() + helper.Check(err) + return result, true +} + +// SaveMetaData stores the metadata of a file to the disk +func (p DatabaseProvider) SaveMetaData(file models.File) { + newData := schemaMetaData{ + Id: file.Id, + Name: file.Name, + Size: file.Size, + SHA1: file.SHA1, + ExpireAt: file.ExpireAt, + SizeBytes: file.SizeBytes, + DownloadsRemaining: file.DownloadsRemaining, + DownloadCount: file.DownloadCount, + PasswordHash: file.PasswordHash, + HotlinkId: file.HotlinkId, + ContentType: file.ContentType, + AwsBucket: file.AwsBucket, + UserId: file.UserId, + UploadDate: file.UploadDate, + PendingDeletion: file.PendingDeletion, + UploadRequestId: file.UploadRequestId, + } + + if file.UnlimitedDownloads { + newData.UnlimitedDownloads = 1 + } + if file.UnlimitedTime { + newData.UnlimitedTime = 1 + } + + var buf bytes.Buffer + enc := gob.NewEncoder(&buf) + err := enc.Encode(file.Encryption) + helper.Check(err) + newData.Encryption = buf.Bytes() + + _, err = p.sqlDb.Exec(`INSERT INTO filemetadata (id, name, size, sha1, expireat, sizebytes, + downloadsremaining, downloadcount, passwordhash, hotlinkid, contenttype, awsbucket, encryption, + unlimiteddownloads, unlimitedtime, userid, uploaddate, pendingdeletion, uploadrequestid) + VALUES ($1, $2, $3, $4, $5, $6, $7, $8, $9, $10, $11, $12, $13, $14, $15, $16, $17, $18, $19) + ON CONFLICT (id) DO UPDATE SET name = $2, size = $3, sha1 = $4, expireat = $5, sizebytes = $6, + downloadsremaining = $7, downloadcount = $8, passwordhash = $9, hotlinkid = $10, contenttype = $11, + awsbucket = $12, encryption = $13, unlimiteddownloads = $14, unlimitedtime = $15, userid = $16, + uploaddate = $17, pendingdeletion = $18, uploadrequestid = $19`, + newData.Id, newData.Name, newData.Size, newData.SHA1, newData.ExpireAt, newData.SizeBytes, + newData.DownloadsRemaining, newData.DownloadCount, newData.PasswordHash, newData.HotlinkId, newData.ContentType, + newData.AwsBucket, newData.Encryption, newData.UnlimitedDownloads, newData.UnlimitedTime, newData.UserId, newData.UploadDate, + newData.PendingDeletion, newData.UploadRequestId) + helper.Check(err) +} + +// IncreaseDownloadCount increases the download count of a file atomically +func (p DatabaseProvider) IncreaseDownloadCount(id string, decreaseRemainingDownloads bool) { + if decreaseRemainingDownloads { + _, err := p.sqlDb.Exec(`UPDATE filemetadata SET downloadcount = downloadcount + 1, + downloadsremaining = downloadsremaining - 1 WHERE id = $1`, id) + helper.Check(err) + } else { + _, err := p.sqlDb.Exec(`UPDATE filemetadata SET downloadcount = downloadcount + 1 WHERE id = $1`, id) + helper.Check(err) + } +} + +// GetDownloadsRemaining returns the remaining downloads of a file that does not implement UnlimitedDownloads +func (p DatabaseProvider) GetDownloadsRemaining(id string) int { + var downloadsRemaining int + row := p.sqlDb.QueryRow("SELECT downloadsremaining FROM filemetadata WHERE id = $1", id) + err := row.Scan(&downloadsRemaining) + if err != nil { + if errors.Is(err, sql.ErrNoRows) { + return 0 + } + helper.Check(err) + return downloadsRemaining + } + return downloadsRemaining +} + +// DeleteMetaData deletes information about a file +func (p DatabaseProvider) DeleteMetaData(id string) { + _, err := p.sqlDb.Exec("DELETE FROM filemetadata WHERE id = $1", id) + helper.Check(err) +} diff --git a/internal/configuration/database/provider/postgres/sessions.go b/internal/configuration/database/provider/postgres/sessions.go new file mode 100644 index 00000000..67259494 --- /dev/null +++ b/internal/configuration/database/provider/postgres/sessions.go @@ -0,0 +1,76 @@ +package postgres + +import ( + "database/sql" + "errors" + "time" + + "github.com/forceu/gokapi/internal/helper" + "github.com/forceu/gokapi/internal/models" +) + +type schemaSessions struct { + Id string + RenewAt int64 + ValidUntil int64 + UserId int +} + +// GetSession returns the session with the given ID or false if not a valid ID +func (p DatabaseProvider) GetSession(id string) (models.Session, bool) { + var rowResult schemaSessions + row := p.sqlDb.QueryRow("SELECT * FROM sessions WHERE id = $1", id) + err := row.Scan(&rowResult.Id, &rowResult.RenewAt, &rowResult.ValidUntil, &rowResult.UserId) + if err != nil { + if errors.Is(err, sql.ErrNoRows) { + return models.Session{}, false + } + helper.Check(err) + return models.Session{}, false + } + result := models.Session{ + RenewAt: rowResult.RenewAt, + ValidUntil: rowResult.ValidUntil, + UserId: rowResult.UserId, + } + return result, true +} + +// SaveSession stores the given session. After the expiry passed, it will be deleted automatically +func (p DatabaseProvider) SaveSession(id string, session models.Session) { + newData := schemaSessions{ + Id: id, + RenewAt: session.RenewAt, + ValidUntil: session.ValidUntil, + UserId: session.UserId, + } + + _, err := p.sqlDb.Exec(`INSERT INTO sessions (id, renewat, validuntil, userid) VALUES ($1, $2, $3, $4) + ON CONFLICT (id) DO UPDATE SET renewat = $2, validuntil = $3, userid = $4`, + newData.Id, newData.RenewAt, newData.ValidUntil, newData.UserId) + helper.Check(err) +} + +// DeleteSession deletes a session with the given ID +func (p DatabaseProvider) DeleteSession(id string) { + _, err := p.sqlDb.Exec("DELETE FROM sessions WHERE id = $1", id) + helper.Check(err) +} + +// DeleteAllSessions logs all users out +func (p DatabaseProvider) DeleteAllSessions() { + //goland:noinspection SqlWithoutWhere + _, err := p.sqlDb.Exec("DELETE FROM sessions") + helper.Check(err) +} + +// DeleteAllSessionsByUser logs the specific users out +func (p DatabaseProvider) DeleteAllSessionsByUser(userId int) { + _, err := p.sqlDb.Exec("DELETE FROM sessions WHERE userid = $1", userId) + helper.Check(err) +} + +func (p DatabaseProvider) cleanExpiredSessions() { + _, err := p.sqlDb.Exec("DELETE FROM sessions WHERE sessions.validuntil < $1", time.Now().Unix()) + helper.Check(err) +} diff --git a/internal/configuration/database/provider/postgres/statistics.go b/internal/configuration/database/provider/postgres/statistics.go new file mode 100644 index 00000000..f2614cb6 --- /dev/null +++ b/internal/configuration/database/provider/postgres/statistics.go @@ -0,0 +1,55 @@ +package postgres + +import ( + "database/sql" + "errors" + + "github.com/forceu/gokapi/internal/helper" +) + +const statIdTraffic = 1 +const statIdTrafficSince = 2 + +// GetStatTraffic returns the total traffic from statistics +func (p DatabaseProvider) GetStatTraffic() uint64 { + var result uint64 + row := p.sqlDb.QueryRow("SELECT value FROM statistics WHERE type = $1", statIdTraffic) + err := row.Scan(&result) + if err != nil { + if errors.Is(err, sql.ErrNoRows) { + return 0 + } + helper.Check(err) + return 0 + } + return result +} + +// SaveStatTraffic stores the total traffic +func (p DatabaseProvider) SaveStatTraffic(totalTraffic uint64) { + _, err := p.sqlDb.Exec(`INSERT INTO statistics (type, value) VALUES ($1, $2) + ON CONFLICT (type) DO UPDATE SET value = $2`, statIdTraffic, totalTraffic) + helper.Check(err) +} + +// SaveTrafficSince stores the beginning of traffic counting +func (p DatabaseProvider) SaveTrafficSince(since int64) { + _, err := p.sqlDb.Exec(`INSERT INTO statistics (type, value) VALUES ($1, $2) + ON CONFLICT (type) DO UPDATE SET value = $2`, statIdTrafficSince, since) + helper.Check(err) +} + +// GetTrafficSince gets the beginning of traffic counting +func (p DatabaseProvider) GetTrafficSince() (int64, bool) { + var result int64 + row := p.sqlDb.QueryRow("SELECT value FROM statistics WHERE type = $1", statIdTrafficSince) + err := row.Scan(&result) + if err != nil { + if errors.Is(err, sql.ErrNoRows) { + return 0, false + } + helper.Check(err) + return 0, false + } + return result, true +} diff --git a/internal/configuration/database/provider/postgres/users.go b/internal/configuration/database/provider/postgres/users.go new file mode 100644 index 00000000..b782d215 --- /dev/null +++ b/internal/configuration/database/provider/postgres/users.go @@ -0,0 +1,112 @@ +package postgres + +import ( + "database/sql" + "errors" + "time" + + "github.com/forceu/gokapi/internal/helper" + "github.com/forceu/gokapi/internal/models" +) + +type schemaUser struct { + Id int + Name string + Password sql.NullString + Permissions models.UserPermission + UserLevel models.UserRank + LastOnline int64 + ResetPassword int +} + +func (s schemaUser) ToUser() models.User { + pw := "" + if s.Password.Valid { + pw = s.Password.String + } + return models.User{ + Id: s.Id, + Name: s.Name, + Permissions: s.Permissions, + UserLevel: s.UserLevel, + LastOnline: s.LastOnline, + Password: pw, + ResetPassword: s.ResetPassword == 1, + } +} + +// GetAllUsers returns a map with all users +func (p DatabaseProvider) GetAllUsers() []models.User { + var result []models.User + rows, err := p.sqlDb.Query("SELECT * FROM users ORDER BY userlevel, lastonline DESC, name") + helper.Check(err) + defer rows.Close() + for rows.Next() { + row := schemaUser{} + err = rows.Scan(&row.Id, &row.Name, &row.Password, &row.Permissions, &row.UserLevel, &row.LastOnline, &row.ResetPassword) + helper.Check(err) + result = append(result, row.ToUser()) + } + return result +} + +func (p DatabaseProvider) getUserWithConstraint(isName bool, searchValue any) (models.User, bool) { + rowResult := schemaUser{} + query := "SELECT * FROM users WHERE id = $1" + if isName { + query = "SELECT * FROM users WHERE name = $1" + } + row := p.sqlDb.QueryRow(query, searchValue) + err := row.Scan(&rowResult.Id, &rowResult.Name, &rowResult.Password, &rowResult.Permissions, &rowResult.UserLevel, &rowResult.LastOnline, &rowResult.ResetPassword) + if err != nil { + if errors.Is(err, sql.ErrNoRows) { + return models.User{}, false + } + helper.Check(err) + return models.User{}, false + } + user := rowResult.ToUser() + return user, true +} + +// GetUser returns a models.User if valid or false if the ID is not valid +func (p DatabaseProvider) GetUser(id int) (models.User, bool) { + return p.getUserWithConstraint(false, id) +} + +// GetUserByName returns a models.User if valid or false if the name is not valid +func (p DatabaseProvider) GetUserByName(username string) (models.User, bool) { + return p.getUserWithConstraint(true, username) +} + +// SaveUser saves a user to the database. If isNewUser is true, a new Id will be generated +func (p DatabaseProvider) SaveUser(user models.User, isNewUser bool) { + resetpw := 0 + if user.ResetPassword { + resetpw = 1 + } + if isNewUser { + _, err := p.sqlDb.Exec("INSERT INTO users (name, password, permissions, userlevel, lastonline, resetpassword) VALUES ($1, $2, $3, $4, $5, $6)", + user.Name, user.Password, user.Permissions, user.UserLevel, user.LastOnline, resetpw) + helper.Check(err) + } else { + _, err := p.sqlDb.Exec(`INSERT INTO users (id, name, password, permissions, userlevel, lastonline, resetpassword) + VALUES ($1, $2, $3, $4, $5, $6, $7) + ON CONFLICT (id) DO UPDATE SET name = $2, password = $3, permissions = $4, userlevel = $5, lastonline = $6, resetpassword = $7`, + user.Id, user.Name, user.Password, user.Permissions, user.UserLevel, user.LastOnline, resetpw) + helper.Check(err) + } +} + +// UpdateUserLastOnline writes the last online time to the database +func (p DatabaseProvider) UpdateUserLastOnline(id int) { + timeNow := time.Now().Unix() + _, err := p.sqlDb.Exec("UPDATE users SET lastonline = $1 WHERE id = $2", timeNow, id) + helper.Check(err) +} + +// DeleteUser deletes a user with the given ID +func (p DatabaseProvider) DeleteUser(id int) { + _, err := p.sqlDb.Exec("DELETE FROM users WHERE id = $1", id) + helper.Check(err) +} diff --git a/internal/configuration/setup/Setup.go b/internal/configuration/setup/Setup.go index 16b19f6c..fb3385fb 100644 --- a/internal/configuration/setup/Setup.go +++ b/internal/configuration/setup/Setup.go @@ -375,6 +375,60 @@ func parseDatabaseSettings(result *models.Configuration, formObjects *[]jsonForm dbUrl.RawQuery = query.Encode() result.DatabaseUrl = dbUrl.String() return nil + case dbabstraction.TypeMariaDb: + host, err := getFormValueString(formObjects, "mariadb_location") + if err != nil { + return err + } + dbName, err := getFormValueString(formObjects, "mariadb_dbname") + if err != nil { + return err + } + mUser, err := getFormValueString(formObjects, "mariadb_user") + if err != nil { + return err + } + mPassword, err := getFormValueString(formObjects, "mariadb_password") + if err != nil { + return err + } + dbUrl := url.URL{ + Scheme: "mariadb", + Host: host, + Path: "/" + dbName, + } + if mUser != "" || mPassword != "" { + dbUrl.User = url.UserPassword(mUser, mPassword) + } + result.DatabaseUrl = dbUrl.String() + return nil + case dbabstraction.TypePostgres: + host, err := getFormValueString(formObjects, "postgres_location") + if err != nil { + return err + } + dbName, err := getFormValueString(formObjects, "postgres_dbname") + if err != nil { + return err + } + pUser, err := getFormValueString(formObjects, "postgres_user") + if err != nil { + return err + } + pPassword, err := getFormValueString(formObjects, "postgres_password") + if err != nil { + return err + } + dbUrl := url.URL{ + Scheme: "postgres", + Host: host, + Path: "/" + dbName, + } + if pUser != "" || pPassword != "" { + dbUrl.User = url.UserPassword(pUser, pPassword) + } + result.DatabaseUrl = dbUrl.String() + return nil default: return errors.New("unsupported database selected") } @@ -383,7 +437,9 @@ func parseDatabaseSettings(result *models.Configuration, formObjects *[]jsonForm // checkForAllDbValues tests if all values were passed, even if they were not required for this particular database // This is done to ensure that no invalid form was passed and makes testing easier func checkForAllDbValues(formObjects *[]jsonFormObject) error { - expectedValues := []string{"dbtype_sel", "sqlite_location", "redis_location", "redis_prefix", "redis_user", "redis_password"} + expectedValues := []string{"dbtype_sel", "sqlite_location", "redis_location", "redis_prefix", "redis_user", "redis_password", + "mariadb_location", "mariadb_dbname", "mariadb_user", "mariadb_password", + "postgres_location", "postgres_dbname", "postgres_user", "postgres_password"} for _, value := range expectedValues { _, err := getFormValueString(formObjects, value) if err != nil { diff --git a/internal/configuration/setup/Setup_test.go b/internal/configuration/setup/Setup_test.go index 8100e7aa..c4566308 100644 --- a/internal/configuration/setup/Setup_test.go +++ b/internal/configuration/setup/Setup_test.go @@ -245,13 +245,21 @@ func TestInitialSetup(t *testing.T) { } type dbFormTest struct { - DatabaseType string `form:"dbtype_sel"` - SqliteLocation string `form:"sqlite_location"` - RedisLocation string `form:"redis_location"` - RedisPrefix string `form:"redis_prefix"` - RedisUser string `form:"redis_user"` - RedisPw string `form:"redis_password"` - RedisUseSsl string `form:"redis_ssl_sel"` + DatabaseType string `form:"dbtype_sel"` + SqliteLocation string `form:"sqlite_location"` + RedisLocation string `form:"redis_location"` + RedisPrefix string `form:"redis_prefix"` + RedisUser string `form:"redis_user"` + RedisPw string `form:"redis_password"` + RedisUseSsl string `form:"redis_ssl_sel"` + MariadbLocation string `form:"mariadb_location"` + MariadbDbName string `form:"mariadb_dbname"` + MariadbUser string `form:"mariadb_user"` + MariadbPw string `form:"mariadb_password"` + PostgresLocation string `form:"postgres_location"` + PostgresDbName string `form:"postgres_dbname"` + PostgresUser string `form:"postgres_user"` + PostgresPw string `form:"postgres_password"` } func generateDbFormValues(input dbFormTest) []jsonFormObject { @@ -301,6 +309,32 @@ func TestParseDatabaseSettings(t *testing.T) { err = parseDatabaseSettings(&output, &input) test.IsNil(t, err) test.IsEqualString(t, output.DatabaseUrl, expected) + + input = generateDbFormValues(dbFormTest{ + DatabaseType: "2", + MariadbLocation: "127.0.0.1:3306", + MariadbDbName: "gokapi", + MariadbUser: "testuser", + MariadbPw: "testpw", + RedisUseSsl: "0", + }) + expected = "mariadb://testuser:testpw@127.0.0.1:3306/gokapi" + err = parseDatabaseSettings(&output, &input) + test.IsNil(t, err) + test.IsEqualString(t, output.DatabaseUrl, expected) + + input = generateDbFormValues(dbFormTest{ + DatabaseType: "3", + PostgresLocation: "127.0.0.1:5432", + PostgresDbName: "gokapi", + PostgresUser: "testuser", + PostgresPw: "testpw", + RedisUseSsl: "0", + }) + expected = "postgres://testuser:testpw@127.0.0.1:5432/gokapi" + err = parseDatabaseSettings(&output, &input) + test.IsNil(t, err) + test.IsEqualString(t, output.DatabaseUrl, expected) } func TestRunConfigModification(t *testing.T) { @@ -566,6 +600,14 @@ type setupValues struct { RedisUser setupEntry `form:"redis_user"` RedisPw setupEntry `form:"redis_password"` RedisUseSsl setupEntry `form:"redis_ssl_sel" isBool:"true"` + MariadbLocation setupEntry `form:"mariadb_location"` + MariadbDbName setupEntry `form:"mariadb_dbname"` + MariadbUser setupEntry `form:"mariadb_user"` + MariadbPw setupEntry `form:"mariadb_password"` + PostgresLocation setupEntry `form:"postgres_location"` + PostgresDbName setupEntry `form:"postgres_dbname"` + PostgresUser setupEntry `form:"postgres_user"` + PostgresPw setupEntry `form:"postgres_password"` } func (s *setupValues) init() { diff --git a/internal/configuration/setup/templates/setup.tmpl b/internal/configuration/setup/templates/setup.tmpl index a42f205c..bffe8f9c 100644 --- a/internal/configuration/setup/templates/setup.tmpl +++ b/internal/configuration/setup/templates/setup.tmpl @@ -118,9 +118,11 @@ {{ end }} - + +
@@ -131,17 +133,37 @@
-
+
-
+
-
+
+ + @@ -800,6 +822,18 @@ function TestAWS(button, isManual) { {{ if eq .DatabaseSettings.Type 0}} document.getElementById("sqlite_location").value = {{.DatabaseSettings.HostUrl}}; + {{ else if eq .DatabaseSettings.Type 2}} + document.getElementById("dbtype_sel").selectedIndex = 2; + document.getElementById("mariadb_location").value = {{.DatabaseSettings.HostUrl}}; + document.getElementById("mariadb_dbname").value = {{.DatabaseSettings.DatabaseName}}; + document.getElementById("mariadb_user").value = {{.DatabaseSettings.Username}}; + document.getElementById("mariadb_password").value = {{.DatabaseSettings.Password}}; + {{ else if eq .DatabaseSettings.Type 3}} + document.getElementById("dbtype_sel").selectedIndex = 3; + document.getElementById("postgres_location").value = {{.DatabaseSettings.HostUrl}}; + document.getElementById("postgres_dbname").value = {{.DatabaseSettings.DatabaseName}}; + document.getElementById("postgres_user").value = {{.DatabaseSettings.Username}}; + document.getElementById("postgres_password").value = {{.DatabaseSettings.Password}}; {{ else }} document.getElementById("dbtype_sel").selectedIndex = 1; document.getElementById("redis_location").value = {{.DatabaseSettings.HostUrl}}; @@ -1094,15 +1128,14 @@ function TestAWS(button, isManual) { function dbChanged() { let divSqlite = document.getElementById("divsqlite"); let divRedis = document.getElementById("divredis"); + let divMariadb = document.getElementById("divmariadb"); + let divPostgres = document.getElementById("divpostgres"); let dbType = document.getElementById("dbtype_sel").value; - - if (dbType == "0") { - divSqlite.style.display = "block"; - divRedis.style.display = "none"; - } else { - divSqlite.style.display = "none"; - divRedis.style.display = "block"; - } + + divSqlite.style.display = dbType == "0" ? "block" : "none"; + divRedis.style.display = dbType == "1" ? "block" : "none"; + divMariadb.style.display = dbType == "2" ? "block" : "none"; + divPostgres.style.display = dbType == "3" ? "block" : "none"; } diff --git a/internal/models/DbConnection.go b/internal/models/DbConnection.go index 26167ea9..dfc8feff 100644 --- a/internal/models/DbConnection.go +++ b/internal/models/DbConnection.go @@ -2,10 +2,11 @@ package models // DbConnection is a struct that contains the database configuration for connecting type DbConnection struct { - HostUrl string - RedisPrefix string - Username string - Password string - RedisUseSsl bool - Type int + HostUrl string + RedisPrefix string + Username string + Password string + DatabaseName string + RedisUseSsl bool + Type int } diff --git a/makefile b/makefile index 6f6d9f49..f1a002ce 100644 --- a/makefile +++ b/makefile @@ -82,13 +82,33 @@ test-specific: .PHONY: test-all test-all: - @echo Testing all tags + @echo Testing all tags @echo go generate ./... go test ./... -parallel 8 --tags=test,noaws go test ./... -parallel 8 --tags=test,awsmock GOKAPI_AWS_BUCKET="gokapi" GOKAPI_AWS_REGION="eu-central-1" GOKAPI_AWS_KEY="keyid" GOKAPI_AWS_KEY_SECRET="secret" go test ./... -parallel 8 --tags=test,awstest +.PHONY: test-mariadb +# Requires a real, reachable MariaDB/MySQL server. Configure with: +# GOKAPI_MARIADB_HOST, GOKAPI_MARIADB_DBNAME, GOKAPI_MARIADB_USER, GOKAPI_MARIADB_PASSWORD +# The target database's tables are dropped and recreated on every run - use a disposable database. +test-mariadb: + @echo Testing MariaDB provider against a real server + @echo + go generate ./... + go test $(GOPACKAGE)/internal/configuration/database/provider/mariadb/... -count=1 -v --tags=test,mariadbtest + +.PHONY: test-postgres +# Requires a real, reachable PostgreSQL server. Configure with: +# GOKAPI_POSTGRES_HOST, GOKAPI_POSTGRES_DBNAME, GOKAPI_POSTGRES_USER, GOKAPI_POSTGRES_PASSWORD +# The target database's tables are dropped and recreated on every run - use a disposable database. +test-postgres: + @echo Testing PostgreSQL provider against a real server + @echo + go generate ./... + go test $(GOPACKAGE)/internal/configuration/database/provider/postgres/... -count=1 -v --tags=test,postgrestest + .PHONY: update-changelog update-changelog: @echo Updaing changelog