mirror of
https://github.com/eitchtee/WYGIWYH.git
synced 2026-09-07 10:27:17 +02:00
Compare commits
| Author | SHA1 | Date | |
|---|---|---|---|
|
|
d500bef481 | ||
|
|
f280fcb172 | ||
|
|
6c808eff38 | ||
|
|
1690a153f9 | ||
|
|
7116176c78 | ||
|
|
d9c9cbe7c3 | ||
|
|
11b03474c1 | ||
|
|
e2f26b3629 | ||
|
|
10c8fcfb97 | ||
|
|
039ad225d3 | ||
|
|
18d4ab7d11 | ||
|
|
c64a363126 | ||
|
|
68a9286ce5 | ||
|
|
1357688e7b | ||
|
|
ba1f421ad3 | ||
|
|
f60f86a5cb | ||
|
|
5588eb33b1 | ||
|
|
934e6bd8e0 | ||
|
|
7933f1c316 | ||
|
|
00f87bea3d | ||
|
|
53ed0d791a | ||
|
|
5147d939ef | ||
|
|
716bd6f57f | ||
|
|
ded4cb37d5 | ||
|
|
cd48edddab | ||
|
|
9e135560fe | ||
|
|
bdbccdec16 | ||
|
|
907c511b59 | ||
|
|
33f0904a0f | ||
|
|
1c4f23a69e | ||
|
|
8373732964 | ||
|
|
8e09f3e8d8 | ||
|
|
1516fe0281 | ||
|
|
3bd16602a0 | ||
|
|
83c2b599d7 | ||
|
|
f9eab3ce17 | ||
|
|
62a33ddd2c | ||
|
|
ea3e9fd68f | ||
|
|
d71bf3ccb3 | ||
|
|
350d78356c | ||
|
|
91dfba519c | ||
|
|
0c9b41182b | ||
|
|
7a900908dd | ||
|
|
dcd52ec189 | ||
|
|
93e73b9aa6 | ||
|
|
7947d5f874 | ||
|
|
2ab6621205 | ||
|
|
39c15eedec | ||
|
|
6dcf5a63f8 | ||
|
|
11a6b3faa6 | ||
|
|
9b1312b744 | ||
|
|
72707f46e6 | ||
|
|
cf908ae87e | ||
|
|
b145d10fbd | ||
|
|
ee39dcf9ef | ||
|
|
5d330ddfbf | ||
|
|
14011c2f60 | ||
|
|
a25adafe3b | ||
|
|
743951a862 | ||
|
|
74d3d5fcf9 | ||
|
|
75f23168ab | ||
|
|
0ec1c5c063 | ||
|
|
e69d5fbd7a | ||
|
|
5a80a3b1d3 | ||
|
|
e77e96879b | ||
|
|
2fbb0221c2 | ||
|
|
d7722ff12c | ||
|
|
83286fff5f | ||
|
|
ad357fac45 | ||
|
|
8d8e87c9b8 | ||
|
|
fee40fd527 | ||
|
|
f59e53f6dc | ||
|
|
2b379987ac | ||
|
|
4b8ccf426d | ||
|
|
4805ce9e04 | ||
|
|
845a8d846b | ||
|
|
1497500c4f | ||
|
|
9e9e60ccec | ||
|
|
ca14f77f41 | ||
|
|
0fb37a59fa | ||
|
|
e74d9177df | ||
|
|
106d721279 | ||
|
|
d0e9c05283 | ||
|
|
4e16831f4d | ||
|
|
7f5a91c11f | ||
|
|
009a7038c8 | ||
|
|
4273c541c5 | ||
|
|
5c4cb16a0a | ||
|
|
9641e169f2 | ||
|
|
25ff0214ab | ||
|
|
0f9d333834 | ||
|
|
ae115cca15 | ||
|
|
bb23ac6df9 | ||
|
|
7db0fcf097 | ||
|
|
02896f21ed | ||
|
|
5082c17d0f | ||
|
|
33570296e0 | ||
|
|
fc99491f78 | ||
|
|
cb0d379261 | ||
|
|
524e390a62 | ||
|
|
db7e22b627 | ||
|
|
6987b54dba | ||
|
|
b44563b09b | ||
|
|
e839f31104 | ||
|
|
2282625790 | ||
|
|
aa8b559152 | ||
|
|
e1862b8241 | ||
|
|
eb6be8548c | ||
|
|
6ee4e21939 | ||
|
|
bdd8aed891 | ||
|
|
801b2a9edd | ||
|
|
968499f1ab | ||
|
|
5b351821b1 | ||
|
|
7b49072848 | ||
|
|
0ee32724f1 | ||
|
|
6a19381672 | ||
|
|
248fec8b4c | ||
|
|
b34c0557fa | ||
|
|
2af4066aab | ||
|
|
d72ff3cdf5 | ||
|
|
63c69e5c6a | ||
|
|
78171183cc | ||
|
|
34a2b6bfd4 | ||
|
|
1dc24f855e | ||
|
|
1390aff07d | ||
|
|
8fc11b0acf | ||
|
|
9a30a0d3c0 | ||
|
|
10eecd09ff | ||
|
|
2cfb3fb12e | ||
|
|
47af8b135b | ||
|
|
39d0e63375 | ||
|
|
792154eba2 | ||
|
|
dc76ed3156 | ||
|
|
e627dd50be | ||
|
|
5527389196 | ||
|
|
be24ca014e | ||
|
|
7c7056536e | ||
|
|
66d5d7a83b | ||
|
|
4d3ce087d6 | ||
|
|
aeaf9fac43 | ||
|
|
dcedf53b83 | ||
|
|
549648bd6b | ||
|
|
79149abdd2 | ||
|
|
d66d1530bb | ||
|
|
4b9b6484d3 | ||
|
|
a02944bdae | ||
|
|
3ede5304f1 | ||
|
|
27041695b8 | ||
|
|
0c927a2fe9 | ||
|
|
c52db80c64 | ||
|
|
330ce8069c | ||
|
|
2989f11b01 | ||
|
|
738bb7fb74 | ||
|
|
79e50cd853 | ||
|
|
0a23c3ad5b | ||
|
|
43c7749102 | ||
|
|
c1c4ccda8c | ||
|
|
615a689c61 | ||
|
|
c5ccc42f99 | ||
|
|
2baa8b21e8 | ||
|
|
2e554141ba | ||
|
|
73ec6dc0fe | ||
|
|
e19449ff99 | ||
|
|
e81651119c | ||
|
|
55e9ef1b3f | ||
|
|
c414179135 | ||
|
|
14c507de0f | ||
|
|
4722690fe9 | ||
|
|
493619a4ff | ||
|
|
fb4aec88f1 | ||
|
|
4a35e770a4 | ||
|
|
83b81edbae | ||
|
|
a7dc2c955e | ||
|
|
ce2ae562c6 | ||
|
|
de2881ffd4 | ||
|
|
838bf22498 | ||
|
|
d3797ae4a5 | ||
|
|
0532397afd | ||
|
|
8106dc58e5 | ||
|
|
5986cf675b | ||
|
|
80da9142f1 | ||
|
|
766516d248 | ||
|
|
3fd0fba1b8 | ||
|
|
c787565c04 | ||
|
|
0413921dbe | ||
|
|
9ecf8279b4 | ||
|
|
86cf625158 | ||
|
|
ea097ab6f0 | ||
|
|
b1201b51bb | ||
|
|
4c1d20215c | ||
|
|
27e85c4776 | ||
|
|
5a73cd20da | ||
|
|
e305fab300 | ||
|
|
c11f525373 | ||
|
|
ea5d86dbf8 | ||
|
|
a1d3539e3c | ||
|
|
1028a11c8b | ||
|
|
e387a5e2a8 | ||
|
|
624dc382cf | ||
|
|
f88699b333 | ||
|
|
ca98dc073b | ||
|
|
63ba7af3c8 | ||
|
|
2d0dee4a9b | ||
|
|
0000a9ee03 | ||
|
|
41adb37fdb | ||
|
|
496651173e | ||
|
|
8836f06b80 | ||
|
|
e98a48b3a7 | ||
|
|
f9bc9f449b | ||
|
|
26eb1ae813 | ||
|
|
29a2cb9813 | ||
|
|
be79e1b25a | ||
|
|
3fd08466a7 | ||
|
|
6896cdcdca | ||
|
|
2532930a64 | ||
|
|
24a1ef2d0a | ||
|
|
163f2f4e5b | ||
|
|
ede63acf5f | ||
|
|
a8ba3d8754 | ||
|
|
e2f1156264 | ||
|
|
d5bbad7887 | ||
|
|
7ebacff6e4 | ||
|
|
df8ef5d04c | ||
|
|
fa2a8b8c65 | ||
|
|
e44ac5dab6 | ||
|
|
f9261d1283 | ||
|
|
4c73c1cae5 | ||
|
|
0315a56f88 | ||
|
|
44d6b8b53c | ||
|
|
3e4d7c6b1f | ||
|
|
63868514f9 | ||
|
|
9055a24327 | ||
|
|
9dc963ed7b | ||
|
|
49cac0588e | ||
|
|
3b2b6d6473 | ||
|
|
db30bcbeb7 | ||
|
|
a122733a47 | ||
|
|
37f3e4d99a | ||
|
|
d756286135 | ||
|
|
06a7378fd8 | ||
|
|
ab4075c500 | ||
|
|
96318f003d | ||
|
|
1a0412264a | ||
|
|
2588404876 | ||
|
|
fdc273103b | ||
|
|
c015b78cd6 | ||
|
|
50e5492ea1 | ||
|
|
796089cdb3 | ||
|
|
c83b1bf2d6 | ||
|
|
b074ef7929 | ||
|
|
ec7e33b3b0 | ||
|
|
72fedea0db | ||
|
|
0a03745ce6 | ||
|
|
ff4bd79634 | ||
|
|
383b42e26d | ||
|
|
48e43ac031 | ||
|
|
21c60c4059 | ||
|
|
dd6a390e6b | ||
|
|
0c961a8250 | ||
|
|
e28c651973 | ||
|
|
7687ff81c3 | ||
|
|
b2d78c9190 | ||
|
|
b0815e00c7 | ||
|
|
fbe9726338 | ||
|
|
0df3a57a33 | ||
|
|
f86613b17a | ||
|
|
ffa4644e1b | ||
|
|
6611559696 | ||
|
|
b455a0251a | ||
|
|
9d7c3212f1 | ||
|
|
0da3185996 | ||
|
|
6c90e1bb7f | ||
|
|
c6543c0841 | ||
|
|
d4740b8406 | ||
|
|
5a51795e6a | ||
|
|
64d7765357 | ||
|
|
070e11ca77 | ||
|
|
39f66b620a | ||
|
|
ad164866e0 | ||
|
|
05c465cb34 | ||
|
|
92cf526b76 | ||
|
|
639236b890 | ||
|
|
519a85d256 | ||
|
|
700d35b5d5 | ||
|
|
10e51971db | ||
|
|
ec0d5fc121 | ||
|
|
01f91352d6 | ||
|
|
63ce57a315 | ||
|
|
eadeb649a1 | ||
|
|
a2871d5289 | ||
|
|
f2a362bc0f | ||
|
|
2076903740 | ||
|
|
c752c0b16e | ||
|
|
1674766253 | ||
|
|
7ea9d56132 | ||
|
|
3699c6c671 | ||
|
|
d7c255aa14 | ||
|
|
d17b9d5736 | ||
|
|
c7ff6db0bf | ||
|
|
a4c7753f69 | ||
|
|
7e08028557 | ||
|
|
5eaf5086d2 | ||
|
|
c949c6cea0 | ||
|
|
71c0e9a271 | ||
|
|
bc65980511 | ||
|
|
ecdb1a52cc | ||
|
|
afc06582b4 | ||
|
|
07cb0a2a0f | ||
|
|
05ede58c36 | ||
|
|
20b6366a18 | ||
|
|
b0101dae1a | ||
|
|
a3d38ff9e0 | ||
|
|
776e2117a0 | ||
|
|
edcad37926 | ||
|
|
2d51d21035 | ||
|
|
94f5c25829 | ||
|
|
88a5c103e5 | ||
|
|
3dce9e1c55 | ||
|
|
41d8564e8b | ||
|
|
5ee2fd244f | ||
|
|
0545fb7651 | ||
|
|
7bd1d2d751 | ||
|
|
9a4ec449df | ||
|
|
f918351303 |
@@ -0,0 +1 @@
|
|||||||
|
__pycache__/
|
||||||
@@ -38,3 +38,21 @@ TASK_WORKERS=1 # This only work if you're using the single container option. Inc
|
|||||||
#OIDC_CLIENT_SECRET=""
|
#OIDC_CLIENT_SECRET=""
|
||||||
#OIDC_SERVER_URL=""
|
#OIDC_SERVER_URL=""
|
||||||
#OIDC_ALLOW_SIGNUP=true
|
#OIDC_ALLOW_SIGNUP=true
|
||||||
|
|
||||||
|
# Personal access tokens. How often (seconds) a token's last_used_at is rewritten.
|
||||||
|
#API_TOKEN_LAST_USED_UPDATE_INTERVAL=600
|
||||||
|
|
||||||
|
# MCP OAuth Application. Uncomment to auto-create/update the OAuth client
|
||||||
|
# used by remote MCP integrations after migrations complete.
|
||||||
|
#MCP_OAUTH_CLIENT_NAME="WYGIWYH MCP"
|
||||||
|
#MCP_OAUTH_CLIENT_ID="mcp-wygiwyh"
|
||||||
|
#MCP_OAUTH_CLIENT_SECRET="<INSERT A SAFE SECRET HERE>"
|
||||||
|
#MCP_OAUTH_REDIRECT_URIS="http://127.0.0.1:8765/callback"
|
||||||
|
#MCP_OAUTH_SKIP_AUTHORIZATION=false
|
||||||
|
|
||||||
|
# Dynamic Client Registration (RFC 7591). Disabled by default because an open
|
||||||
|
# registration endpoint lets anyone create OAuth applications. Enable only if
|
||||||
|
# remote MCP clients must self-register, and optionally require an initial
|
||||||
|
# access token (sent as "Authorization: Bearer <token>" on /oauth/register/).
|
||||||
|
#OAUTH2_DCR_ENABLED=false
|
||||||
|
#OAUTH2_DCR_INITIAL_ACCESS_TOKEN=""
|
||||||
|
|||||||
@@ -32,15 +32,16 @@ jobs:
|
|||||||
token: ${{ secrets.PAT }}
|
token: ${{ secrets.PAT }}
|
||||||
ref: ${{ github.head_ref }}
|
ref: ${{ github.head_ref }}
|
||||||
|
|
||||||
- name: Set up Python 3.11
|
- name: Install uv
|
||||||
uses: actions/setup-python@v4
|
uses: astral-sh/setup-uv@v5
|
||||||
with:
|
with:
|
||||||
python-version: '3.11'
|
enable-cache: true
|
||||||
|
|
||||||
|
- name: Set up Python 3.11
|
||||||
|
run: uv python install 3.11
|
||||||
|
|
||||||
- name: Install dependencies
|
- name: Install dependencies
|
||||||
run: |
|
run: uv sync --frozen --no-dev
|
||||||
python -m pip install --upgrade pip
|
|
||||||
pip install -r requirements.txt
|
|
||||||
|
|
||||||
- name: Install gettext
|
- name: Install gettext
|
||||||
run: sudo apt-get install -y gettext
|
run: sudo apt-get install -y gettext
|
||||||
@@ -48,7 +49,7 @@ jobs:
|
|||||||
- name: Run makemessages
|
- name: Run makemessages
|
||||||
run: |
|
run: |
|
||||||
cd app
|
cd app
|
||||||
python manage.py makemessages -a
|
uv run python manage.py makemessages -a
|
||||||
|
|
||||||
- name: Check for changes
|
- name: Check for changes
|
||||||
id: check_changes
|
id: check_changes
|
||||||
@@ -64,7 +65,6 @@ jobs:
|
|||||||
if: steps.check_changes.outputs.changes_detected == 'true'
|
if: steps.check_changes.outputs.changes_detected == 'true'
|
||||||
uses: stefanzweifel/git-auto-commit-action@v5
|
uses: stefanzweifel/git-auto-commit-action@v5
|
||||||
with:
|
with:
|
||||||
push_options: --force
|
|
||||||
commit_message: |
|
commit_message: |
|
||||||
chore(locale): update translation files
|
chore(locale): update translation files
|
||||||
|
|
||||||
|
|||||||
@@ -165,3 +165,6 @@ cython_debug/
|
|||||||
node_modules/
|
node_modules/
|
||||||
postgres_data/
|
postgres_data/
|
||||||
.prod.env
|
.prod.env
|
||||||
|
|
||||||
|
# Private local uploads
|
||||||
|
app/attachments/
|
||||||
|
|||||||
Vendored
+29
@@ -0,0 +1,29 @@
|
|||||||
|
{
|
||||||
|
"version": "0.2.0",
|
||||||
|
"configurations": [
|
||||||
|
{
|
||||||
|
"name": "Docker: Dev",
|
||||||
|
"type": "node-terminal",
|
||||||
|
"request": "launch",
|
||||||
|
"command": "docker compose --env-file .env -f docker-compose.dev.yml up --build",
|
||||||
|
"cwd": "${workspaceFolder}",
|
||||||
|
"postDebugTask": "Docker: Dev Down"
|
||||||
|
},
|
||||||
|
{
|
||||||
|
"name": "Docker: Dev (no rebuild)",
|
||||||
|
"type": "node-terminal",
|
||||||
|
"request": "launch",
|
||||||
|
"command": "docker compose --env-file .env -f docker-compose.dev.yml up",
|
||||||
|
"cwd": "${workspaceFolder}",
|
||||||
|
"postDebugTask": "Docker: Dev Down"
|
||||||
|
},
|
||||||
|
{
|
||||||
|
"name": "Docker: Prod",
|
||||||
|
"type": "node-terminal",
|
||||||
|
"request": "launch",
|
||||||
|
"command": "docker compose --env-file .prod.env -f docker-compose.prod.yml up --build",
|
||||||
|
"cwd": "${workspaceFolder}",
|
||||||
|
"postDebugTask": "Docker: Prod Down"
|
||||||
|
}
|
||||||
|
]
|
||||||
|
}
|
||||||
Vendored
+119
@@ -0,0 +1,119 @@
|
|||||||
|
{
|
||||||
|
"version": "2.0.0",
|
||||||
|
"tasks": [
|
||||||
|
{
|
||||||
|
"label": "Docker: Dev",
|
||||||
|
"type": "shell",
|
||||||
|
"command": "docker",
|
||||||
|
"args": [
|
||||||
|
"compose",
|
||||||
|
"--env-file",
|
||||||
|
".env",
|
||||||
|
"-f",
|
||||||
|
"docker-compose.dev.yml",
|
||||||
|
"up",
|
||||||
|
"--build"
|
||||||
|
],
|
||||||
|
"options": {
|
||||||
|
"cwd": "${workspaceFolder}"
|
||||||
|
},
|
||||||
|
"group": "build",
|
||||||
|
"problemMatcher": []
|
||||||
|
},
|
||||||
|
{
|
||||||
|
"label": "Docker: Dev (no rebuild)",
|
||||||
|
"type": "shell",
|
||||||
|
"command": "docker",
|
||||||
|
"args": [
|
||||||
|
"compose",
|
||||||
|
"--env-file",
|
||||||
|
".env",
|
||||||
|
"-f",
|
||||||
|
"docker-compose.dev.yml",
|
||||||
|
"up"
|
||||||
|
],
|
||||||
|
"options": {
|
||||||
|
"cwd": "${workspaceFolder}"
|
||||||
|
},
|
||||||
|
"problemMatcher": []
|
||||||
|
},
|
||||||
|
{
|
||||||
|
"label": "Docker: Dev Refresh Vite Deps",
|
||||||
|
"type": "shell",
|
||||||
|
"command": "docker compose --env-file .env -f docker-compose.dev.yml rm -sfv vite; docker compose --env-file .env -f docker-compose.dev.yml up --build",
|
||||||
|
"options": {
|
||||||
|
"cwd": "${workspaceFolder}"
|
||||||
|
},
|
||||||
|
"problemMatcher": []
|
||||||
|
},
|
||||||
|
{
|
||||||
|
"label": "Docker: Dev Down",
|
||||||
|
"type": "shell",
|
||||||
|
"command": "docker",
|
||||||
|
"args": [
|
||||||
|
"compose",
|
||||||
|
"--env-file",
|
||||||
|
".env",
|
||||||
|
"-f",
|
||||||
|
"docker-compose.dev.yml",
|
||||||
|
"down"
|
||||||
|
],
|
||||||
|
"options": {
|
||||||
|
"cwd": "${workspaceFolder}"
|
||||||
|
},
|
||||||
|
"problemMatcher": []
|
||||||
|
},
|
||||||
|
{
|
||||||
|
"label": "Docker: Prod",
|
||||||
|
"type": "shell",
|
||||||
|
"command": "docker",
|
||||||
|
"args": [
|
||||||
|
"compose",
|
||||||
|
"--env-file",
|
||||||
|
".prod.env",
|
||||||
|
"-f",
|
||||||
|
"docker-compose.prod.yml",
|
||||||
|
"up",
|
||||||
|
"--build"
|
||||||
|
],
|
||||||
|
"options": {
|
||||||
|
"cwd": "${workspaceFolder}"
|
||||||
|
},
|
||||||
|
"problemMatcher": []
|
||||||
|
},
|
||||||
|
{
|
||||||
|
"label": "Docker: Prod Down",
|
||||||
|
"type": "shell",
|
||||||
|
"command": "docker",
|
||||||
|
"args": [
|
||||||
|
"compose",
|
||||||
|
"--env-file",
|
||||||
|
".prod.env",
|
||||||
|
"-f",
|
||||||
|
"docker-compose.prod.yml",
|
||||||
|
"down"
|
||||||
|
],
|
||||||
|
"options": {
|
||||||
|
"cwd": "${workspaceFolder}"
|
||||||
|
},
|
||||||
|
"problemMatcher": []
|
||||||
|
},
|
||||||
|
{
|
||||||
|
"label": "Django: Runserver localhost:8000",
|
||||||
|
"type": "shell",
|
||||||
|
"command": "${command:python.interpreterPath}",
|
||||||
|
"args": [
|
||||||
|
"manage.py",
|
||||||
|
"runserver",
|
||||||
|
"localhost:8000"
|
||||||
|
],
|
||||||
|
"options": {
|
||||||
|
"cwd": "${workspaceFolder}/app",
|
||||||
|
"env": {
|
||||||
|
"PYTHONUNBUFFERED": "1"
|
||||||
|
}
|
||||||
|
},
|
||||||
|
"problemMatcher": []
|
||||||
|
}
|
||||||
|
]
|
||||||
|
}
|
||||||
@@ -157,6 +157,13 @@ WYGIWYH supports login via OpenID Connect (OIDC) through `django-allauth`. This
|
|||||||
> [!NOTE]
|
> [!NOTE]
|
||||||
> Currently only OpenID Connect is supported as a provider, open an issue if you need something else.
|
> Currently only OpenID Connect is supported as a provider, open an issue if you need something else.
|
||||||
|
|
||||||
|
> [!Caution]
|
||||||
|
> WYGIWYH automatically connects OIDC accounts to existing local accounts with matching email addresses.
|
||||||
|
> This means if a user already exists with email `user@example.com` and someone logs in via OIDC with the same email, the OIDC account will be automatically linked to the existing account without requiring user confirmation.
|
||||||
|
> This is only recommended for trusted OIDC providers that verify email addresses and where you control who can create accounts.
|
||||||
|
|
||||||
|
### Configuration
|
||||||
|
|
||||||
To configure OIDC, you need to set the following environment variables:
|
To configure OIDC, you need to set the following environment variables:
|
||||||
|
|
||||||
| Variable | Description |
|
| Variable | Description |
|
||||||
@@ -175,6 +182,49 @@ When configuring your OIDC provider, you will need to provide a callback URL (al
|
|||||||
|
|
||||||
Replace `https://your.wygiwyh.domain` with the actual URL where your WYGIWYH instance is accessible. And `<OIDC_CLIENT_NAME>` with the slugfied value set in OIDC_CLIENT_NAME or the default `openid-connect` if you haven't set this variable.
|
Replace `https://your.wygiwyh.domain` with the actual URL where your WYGIWYH instance is accessible. And `<OIDC_CLIENT_NAME>` with the slugfied value set in OIDC_CLIENT_NAME or the default `openid-connect` if you haven't set this variable.
|
||||||
|
|
||||||
|
### API Tokens for n8n and other automations
|
||||||
|
|
||||||
|
If you need a stable non-browser credential for automations such as `n8n`, WYGIWYH can also issue its own user-bound API tokens. This avoids Keycloak login flows and can be used directly against `/api/`.
|
||||||
|
|
||||||
|
Create a token from the container or application shell:
|
||||||
|
|
||||||
|
```bash
|
||||||
|
python manage.py create_api_token you@example.com --name n8n
|
||||||
|
```
|
||||||
|
|
||||||
|
Optional expiration:
|
||||||
|
|
||||||
|
```bash
|
||||||
|
python manage.py create_api_token you@example.com --name n8n --expires-in-days 90
|
||||||
|
```
|
||||||
|
|
||||||
|
The command prints the raw token **once**. Store it in your secret manager and use it like this:
|
||||||
|
|
||||||
|
```bash
|
||||||
|
curl -H "Authorization: Token wygiwyh_pat_<key>.<secret>" \
|
||||||
|
https://your.wygiwyh.domain/api/accounts/
|
||||||
|
```
|
||||||
|
|
||||||
|
Recommended usage for automation is a dedicated WYGIWYH user such as `n8n@...`, so API ownership and audit trails stay separate from your interactive account.
|
||||||
|
|
||||||
|
### MCP OAuth Application Bootstrap
|
||||||
|
|
||||||
|
If you want WYGIWYH to act as the OAuth authorization server for a remote MCP server, you can let the container create or update the OAuth application automatically on startup.
|
||||||
|
|
||||||
|
Set these environment variables:
|
||||||
|
|
||||||
|
| Variable | Description |
|
||||||
|
|---|---|
|
||||||
|
| `MCP_OAUTH_CLIENT_NAME` | Optional display name for the OAuth client. Defaults to `WYGIWYH MCP`. |
|
||||||
|
| `MCP_OAUTH_CLIENT_ID` | Client ID that will be created or updated in `django-oauth-toolkit`. |
|
||||||
|
| `MCP_OAUTH_CLIENT_SECRET` | Client secret for that OAuth application. |
|
||||||
|
| `MCP_OAUTH_REDIRECT_URIS` | Space-separated redirect URIs allowed for the MCP OAuth client. |
|
||||||
|
| `MCP_OAUTH_SKIP_AUTHORIZATION` | Set to `true` to bypass the consent screen. Defaults to `false`. |
|
||||||
|
|
||||||
|
When these variables are present, startup runs `python manage.py setup_oauth` after migrations and keeps the OAuth application in sync without needing a manual Django admin step.
|
||||||
|
|
||||||
|
WYGIWYH also exposes OAuth Dynamic Client Registration at `/.well-known/oauth-authorization-server` via `registration_endpoint`, so MCP clients that support RFC 7591 can self-register instead of relying on a pre-created `MCP_OAUTH_CLIENT_ID` / `MCP_OAUTH_CLIENT_SECRET`. The current implementation supports `authorization_code` + PKCE clients using `none`, `client_secret_basic`, or `client_secret_post` token endpoint auth methods.
|
||||||
|
|
||||||
# How it works
|
# How it works
|
||||||
|
|
||||||
Check out our [Wiki](https://github.com/eitchtee/WYGIWYH/wiki) for more information.
|
Check out our [Wiki](https://github.com/eitchtee/WYGIWYH/wiki) for more information.
|
||||||
|
|||||||
+75
-23
@@ -70,7 +70,9 @@ INSTALLED_APPS = [
|
|||||||
"apps.api.apps.ApiConfig",
|
"apps.api.apps.ApiConfig",
|
||||||
"cachalot",
|
"cachalot",
|
||||||
"rest_framework",
|
"rest_framework",
|
||||||
|
"rest_framework.authtoken",
|
||||||
"drf_spectacular",
|
"drf_spectacular",
|
||||||
|
"oauth2_provider",
|
||||||
"django_cotton",
|
"django_cotton",
|
||||||
"apps.rules.apps.RulesConfig",
|
"apps.rules.apps.RulesConfig",
|
||||||
"apps.calendar_view.apps.CalendarViewConfig",
|
"apps.calendar_view.apps.CalendarViewConfig",
|
||||||
@@ -143,6 +145,9 @@ WSGI_APPLICATION = "WYGIWYH.wsgi.application"
|
|||||||
# Database
|
# Database
|
||||||
# https://docs.djangoproject.com/en/5.1/ref/settings/#databases
|
# https://docs.djangoproject.com/en/5.1/ref/settings/#databases
|
||||||
|
|
||||||
|
THREADS = int(os.getenv("GUNICORN_THREADS", 1))
|
||||||
|
MAX_POOL_SIZE = THREADS + 1
|
||||||
|
|
||||||
DATABASES = {
|
DATABASES = {
|
||||||
"default": {
|
"default": {
|
||||||
"ENGINE": "django.db.backends.postgresql",
|
"ENGINE": "django.db.backends.postgresql",
|
||||||
@@ -151,8 +156,16 @@ DATABASES = {
|
|||||||
"PASSWORD": os.getenv("SQL_PASSWORD", "password"),
|
"PASSWORD": os.getenv("SQL_PASSWORD", "password"),
|
||||||
"HOST": os.getenv("SQL_HOST", "localhost"),
|
"HOST": os.getenv("SQL_HOST", "localhost"),
|
||||||
"PORT": os.getenv("SQL_PORT", "5432"),
|
"PORT": os.getenv("SQL_PORT", "5432"),
|
||||||
|
"CONN_MAX_AGE": 0,
|
||||||
|
"CONN_HEALTH_CHECKS": True,
|
||||||
"OPTIONS": {
|
"OPTIONS": {
|
||||||
"pool": True,
|
"pool": {
|
||||||
|
"min_size": 1,
|
||||||
|
"max_size": MAX_POOL_SIZE,
|
||||||
|
"timeout": 10,
|
||||||
|
"max_lifetime": 600,
|
||||||
|
"max_idle": 300,
|
||||||
|
},
|
||||||
},
|
},
|
||||||
}
|
}
|
||||||
}
|
}
|
||||||
@@ -299,6 +312,7 @@ LOCALE_PATHS = [BASE_DIR / "locale"]
|
|||||||
|
|
||||||
STATIC_URL = "static/"
|
STATIC_URL = "static/"
|
||||||
STATIC_ROOT = BASE_DIR / "static_files"
|
STATIC_ROOT = BASE_DIR / "static_files"
|
||||||
|
ATTACHMENT_MEDIA_ROOT = BASE_DIR / "attachments"
|
||||||
|
|
||||||
STATICFILES_DIRS = [
|
STATICFILES_DIRS = [
|
||||||
ROOT_DIR / "frontend" / "build",
|
ROOT_DIR / "frontend" / "build",
|
||||||
@@ -331,6 +345,11 @@ DEFAULT_AUTO_FIELD = "django.db.models.BigAutoField"
|
|||||||
LOGIN_REDIRECT_URL = "/"
|
LOGIN_REDIRECT_URL = "/"
|
||||||
LOGIN_URL = "/login/"
|
LOGIN_URL = "/login/"
|
||||||
LOGOUT_REDIRECT_URL = "/login/"
|
LOGOUT_REDIRECT_URL = "/login/"
|
||||||
|
# Public base URL advertised in OAuth metadata. Falls back to the first entry
|
||||||
|
# of the existing space-separated URL env var, then to the request host.
|
||||||
|
PUBLIC_BASE_URL = (
|
||||||
|
os.getenv("PUBLIC_BASE_URL", "") or os.getenv("URL", "").split(" ")[0]
|
||||||
|
).rstrip("/")
|
||||||
|
|
||||||
# Allauth settings
|
# Allauth settings
|
||||||
AUTHENTICATION_BACKENDS = [
|
AUTHENTICATION_BACKENDS = [
|
||||||
@@ -364,8 +383,16 @@ ACCOUNT_EMAIL_VERIFICATION = "none"
|
|||||||
SOCIALACCOUNT_LOGIN_ON_GET = True
|
SOCIALACCOUNT_LOGIN_ON_GET = True
|
||||||
SOCIALACCOUNT_ONLY = True
|
SOCIALACCOUNT_ONLY = True
|
||||||
SOCIALACCOUNT_AUTO_SIGNUP = os.getenv("OIDC_ALLOW_SIGNUP", "true").lower() == "true"
|
SOCIALACCOUNT_AUTO_SIGNUP = os.getenv("OIDC_ALLOW_SIGNUP", "true").lower() == "true"
|
||||||
|
SOCIALACCOUNT_EMAIL_AUTHENTICATION = True
|
||||||
|
SOCIALACCOUNT_EMAIL_AUTHENTICATION_AUTO_CONNECT = True
|
||||||
ACCOUNT_ADAPTER = "allauth.account.adapter.DefaultAccountAdapter"
|
ACCOUNT_ADAPTER = "allauth.account.adapter.DefaultAccountAdapter"
|
||||||
SOCIALACCOUNT_ADAPTER = "allauth.socialaccount.adapter.DefaultSocialAccountAdapter"
|
SOCIALACCOUNT_ADAPTER = "apps.users.adapters.AutoConnectSocialAccountAdapter"
|
||||||
|
|
||||||
|
# Personal access tokens. last_used_at is only rewritten once per interval to
|
||||||
|
# avoid a database write on every authenticated request.
|
||||||
|
API_TOKEN_LAST_USED_UPDATE_INTERVAL = int(
|
||||||
|
os.getenv("API_TOKEN_LAST_USED_UPDATE_INTERVAL", "600")
|
||||||
|
)
|
||||||
|
|
||||||
# CRISPY FORMS
|
# CRISPY FORMS
|
||||||
CRISPY_ALLOWED_TEMPLATE_PACKS = [
|
CRISPY_ALLOWED_TEMPLATE_PACKS = [
|
||||||
@@ -378,6 +405,10 @@ SESSION_EXPIRE_AT_BROWSER_CLOSE = False
|
|||||||
SESSION_COOKIE_AGE = int(os.getenv("SESSION_EXPIRY_TIME", 2678400)) # 31 days
|
SESSION_COOKIE_AGE = int(os.getenv("SESSION_EXPIRY_TIME", 2678400)) # 31 days
|
||||||
SESSION_COOKIE_SECURE = os.getenv("HTTPS_ENABLED", "false").lower() == "true"
|
SESSION_COOKIE_SECURE = os.getenv("HTTPS_ENABLED", "false").lower() == "true"
|
||||||
|
|
||||||
|
HTTPS_ENABLED = os.getenv("HTTPS_ENABLED", "false").lower() == "true"
|
||||||
|
ACCOUNT_DEFAULT_HTTP_PROTOCOL = "https" if HTTPS_ENABLED else "http"
|
||||||
|
SECURE_PROXY_SSL_HEADER = ("HTTP_X_FORWARDED_PROTO", "https") if HTTPS_ENABLED else None
|
||||||
|
|
||||||
DEBUG_TOOLBAR_CONFIG = {
|
DEBUG_TOOLBAR_CONFIG = {
|
||||||
"ROOT_TAG_EXTRA_ATTRS": "hx-preserve",
|
"ROOT_TAG_EXTRA_ATTRS": "hx-preserve",
|
||||||
# "SHOW_TOOLBAR_CALLBACK": lambda r: False, # disables it
|
# "SHOW_TOOLBAR_CALLBACK": lambda r: False, # disables it
|
||||||
@@ -422,11 +453,38 @@ REST_FRAMEWORK = {
|
|||||||
"apps.api.permissions.NotInDemoMode",
|
"apps.api.permissions.NotInDemoMode",
|
||||||
"rest_framework.permissions.DjangoModelPermissions",
|
"rest_framework.permissions.DjangoModelPermissions",
|
||||||
],
|
],
|
||||||
"DEFAULT_PAGINATION_CLASS": "rest_framework.pagination.PageNumberPagination",
|
"DEFAULT_FILTER_BACKENDS": [
|
||||||
"PAGE_SIZE": 10,
|
"django_filters.rest_framework.DjangoFilterBackend",
|
||||||
|
"rest_framework.filters.OrderingFilter",
|
||||||
|
],
|
||||||
|
"DEFAULT_AUTHENTICATION_CLASSES": [
|
||||||
|
"oauth2_provider.contrib.rest_framework.OAuth2Authentication",
|
||||||
|
"apps.api.authentication.APITokenAuthentication",
|
||||||
|
"rest_framework.authentication.TokenAuthentication",
|
||||||
|
"rest_framework.authentication.SessionAuthentication",
|
||||||
|
"rest_framework.authentication.BasicAuthentication",
|
||||||
|
],
|
||||||
|
"DEFAULT_PAGINATION_CLASS": "apps.api.custom.pagination.CustomPageNumberPagination",
|
||||||
"DEFAULT_SCHEMA_CLASS": "drf_spectacular.openapi.AutoSchema",
|
"DEFAULT_SCHEMA_CLASS": "drf_spectacular.openapi.AutoSchema",
|
||||||
}
|
}
|
||||||
|
|
||||||
|
OAUTH2_PROVIDER = {
|
||||||
|
"PKCE_REQUIRED": True,
|
||||||
|
"ACCESS_TOKEN_EXPIRE_SECONDS": int(
|
||||||
|
os.getenv("OAUTH2_ACCESS_TOKEN_EXPIRE_SECONDS", "3600")
|
||||||
|
),
|
||||||
|
"SCOPES": {
|
||||||
|
"mcp": "Access WYGIWYH from MCP clients.",
|
||||||
|
},
|
||||||
|
}
|
||||||
|
|
||||||
|
# Dynamic Client Registration (RFC 7591). Disabled by default: an open
|
||||||
|
# registration endpoint lets anyone create OAuth applications. Enable it only
|
||||||
|
# when remote MCP clients must self-register, and optionally require an initial
|
||||||
|
# access token presented as `Authorization: Bearer <token>`.
|
||||||
|
OAUTH2_DCR_ENABLED = os.getenv("OAUTH2_DCR_ENABLED", "false").lower() == "true"
|
||||||
|
OAUTH2_DCR_INITIAL_ACCESS_TOKEN = os.getenv("OAUTH2_DCR_INITIAL_ACCESS_TOKEN", "")
|
||||||
|
|
||||||
SPECTACULAR_SETTINGS = {
|
SPECTACULAR_SETTINGS = {
|
||||||
"TITLE": "WYGIWYH API",
|
"TITLE": "WYGIWYH API",
|
||||||
"DESCRIPTION": "A no-frills expense tracker",
|
"DESCRIPTION": "A no-frills expense tracker",
|
||||||
@@ -438,7 +496,7 @@ SPECTACULAR_SETTINGS = {
|
|||||||
if "procrastinate" in sys.argv:
|
if "procrastinate" in sys.argv:
|
||||||
LOGGING = {
|
LOGGING = {
|
||||||
"version": 1,
|
"version": 1,
|
||||||
"disable_existing_loggers": False,
|
"disable_existing_loggers": True,
|
||||||
"formatters": {
|
"formatters": {
|
||||||
"standard": {
|
"standard": {
|
||||||
"format": "[%(asctime)s] - %(levelname)s - %(name)s - %(message)s",
|
"format": "[%(asctime)s] - %(levelname)s - %(name)s - %(message)s",
|
||||||
@@ -446,26 +504,19 @@ if "procrastinate" in sys.argv:
|
|||||||
},
|
},
|
||||||
},
|
},
|
||||||
"handlers": {
|
"handlers": {
|
||||||
"procrastinate": {
|
|
||||||
"level": "INFO",
|
|
||||||
"class": "logging.StreamHandler",
|
|
||||||
"formatter": "standard",
|
|
||||||
},
|
|
||||||
"console": {
|
"console": {
|
||||||
"class": "logging.StreamHandler",
|
"class": "logging.StreamHandler",
|
||||||
"formatter": "standard",
|
"formatter": "standard",
|
||||||
"level": "INFO",
|
"level": "INFO",
|
||||||
},
|
},
|
||||||
},
|
},
|
||||||
|
"root": {
|
||||||
|
"handlers": ["console"],
|
||||||
|
"level": "INFO",
|
||||||
|
},
|
||||||
"loggers": {
|
"loggers": {
|
||||||
"procrastinate": {
|
"procrastinate": {
|
||||||
"handlers": ["procrastinate"],
|
|
||||||
"propagate": False,
|
|
||||||
},
|
|
||||||
"root": {
|
|
||||||
"handlers": ["console"],
|
|
||||||
"level": "INFO",
|
"level": "INFO",
|
||||||
"propagate": False,
|
|
||||||
},
|
},
|
||||||
},
|
},
|
||||||
}
|
}
|
||||||
@@ -485,19 +536,20 @@ else:
|
|||||||
"formatter": "standard",
|
"formatter": "standard",
|
||||||
"level": "INFO",
|
"level": "INFO",
|
||||||
},
|
},
|
||||||
"procrastinate": {
|
},
|
||||||
"level": "INFO",
|
"root": {
|
||||||
"class": "logging.StreamHandler",
|
"handlers": ["console"],
|
||||||
},
|
"level": "INFO",
|
||||||
},
|
},
|
||||||
"loggers": {
|
"loggers": {
|
||||||
"procrastinate": {
|
"procrastinate": {
|
||||||
"handlers": None,
|
"handlers": [],
|
||||||
"propagate": False,
|
"propagate": False,
|
||||||
},
|
},
|
||||||
"root": {
|
"allauth": {
|
||||||
"handlers": ["console"],
|
"handlers": ["console"],
|
||||||
"level": "INFO",
|
"level": "DEBUG",
|
||||||
|
"propagate": False,
|
||||||
},
|
},
|
||||||
},
|
},
|
||||||
}
|
}
|
||||||
|
|||||||
@@ -22,6 +22,29 @@ from drf_spectacular.views import (
|
|||||||
SpectacularSwaggerView,
|
SpectacularSwaggerView,
|
||||||
)
|
)
|
||||||
from allauth.socialaccount.providers.openid_connect.views import login, callback
|
from allauth.socialaccount.providers.openid_connect.views import login, callback
|
||||||
|
from apps.common.decorators.demo import disabled_on_demo
|
||||||
|
from apps.common.oauth_views import (
|
||||||
|
authorization_server_metadata,
|
||||||
|
dynamic_client_registration,
|
||||||
|
)
|
||||||
|
from oauth2_provider import urls as _dot_urls
|
||||||
|
|
||||||
|
|
||||||
|
def _decorate_included(patterns, decorator):
|
||||||
|
"""Apply ``decorator`` to every view callback inside an included URLconf.
|
||||||
|
|
||||||
|
django.urls does not support decorating ``include()`` directly, so we wrap
|
||||||
|
each URLPattern's callback here. The OAuth2 endpoints issue credentials, so
|
||||||
|
gate them behind the same DEMO-mode guard used elsewhere.
|
||||||
|
"""
|
||||||
|
wrapped = []
|
||||||
|
for pattern in patterns:
|
||||||
|
pattern.callback = decorator(pattern.callback)
|
||||||
|
wrapped.append(pattern)
|
||||||
|
return wrapped
|
||||||
|
|
||||||
|
|
||||||
|
_oauth_patterns = _decorate_included(_dot_urls.urlpatterns, disabled_on_demo)
|
||||||
|
|
||||||
|
|
||||||
urlpatterns = [
|
urlpatterns = [
|
||||||
@@ -39,6 +62,20 @@ urlpatterns = [
|
|||||||
name="swagger-ui",
|
name="swagger-ui",
|
||||||
),
|
),
|
||||||
path("auth/", include("allauth.urls")), # allauth urls
|
path("auth/", include("allauth.urls")), # allauth urls
|
||||||
|
path(
|
||||||
|
"oauth/",
|
||||||
|
include((_oauth_patterns, _dot_urls.app_name), namespace="oauth2_provider"),
|
||||||
|
),
|
||||||
|
path(
|
||||||
|
".well-known/oauth-authorization-server",
|
||||||
|
disabled_on_demo(authorization_server_metadata),
|
||||||
|
name="oauth-authorization-server-metadata",
|
||||||
|
),
|
||||||
|
path(
|
||||||
|
"oauth/register/",
|
||||||
|
disabled_on_demo(dynamic_client_registration),
|
||||||
|
name="oauth-dynamic-client-registration",
|
||||||
|
),
|
||||||
# path("auth/oidc/<str:provider_id>/login/", login, name="openid_connect_login"),
|
# path("auth/oidc/<str:provider_id>/login/", login, name="openid_connect_login"),
|
||||||
# path(
|
# path(
|
||||||
# "auth/oidc/<str:provider_id>/login/callback/",
|
# "auth/oidc/<str:provider_id>/login/callback/",
|
||||||
|
|||||||
@@ -1,13 +1,19 @@
|
|||||||
from django.contrib import messages
|
from django.contrib import messages
|
||||||
from django.contrib.auth.decorators import login_required
|
from django.contrib.auth.decorators import login_required
|
||||||
|
from django.core.exceptions import PermissionDenied
|
||||||
from django.http import HttpResponse
|
from django.http import HttpResponse
|
||||||
from django.shortcuts import render, get_object_or_404
|
from django.shortcuts import render
|
||||||
from django.utils.translation import gettext_lazy as _
|
from django.utils.translation import gettext_lazy as _
|
||||||
from django.views.decorators.http import require_http_methods
|
from django.views.decorators.http import require_http_methods
|
||||||
|
|
||||||
from apps.accounts.forms import AccountGroupForm
|
from apps.accounts.forms import AccountGroupForm
|
||||||
from apps.accounts.models import AccountGroup
|
from apps.accounts.models import AccountGroup
|
||||||
from apps.common.decorators.htmx import only_htmx
|
from apps.common.decorators.htmx import only_htmx
|
||||||
|
from apps.common.functions.permissions import (
|
||||||
|
EDIT,
|
||||||
|
READ,
|
||||||
|
get_shared_object_or_error,
|
||||||
|
)
|
||||||
from apps.common.models import SharedObject
|
from apps.common.models import SharedObject
|
||||||
from apps.common.forms import SharedObjectForm
|
from apps.common.forms import SharedObjectForm
|
||||||
|
|
||||||
@@ -25,7 +31,7 @@ def account_groups_index(request):
|
|||||||
@login_required
|
@login_required
|
||||||
@require_http_methods(["GET"])
|
@require_http_methods(["GET"])
|
||||||
def account_groups_list(request):
|
def account_groups_list(request):
|
||||||
account_groups = AccountGroup.objects.all().order_by("id")
|
account_groups = AccountGroup.objects.all().order_by("name")
|
||||||
return render(
|
return render(
|
||||||
request,
|
request,
|
||||||
"account_groups/fragments/list.html",
|
"account_groups/fragments/list.html",
|
||||||
@@ -63,17 +69,7 @@ def account_group_add(request, **kwargs):
|
|||||||
@login_required
|
@login_required
|
||||||
@require_http_methods(["GET", "POST"])
|
@require_http_methods(["GET", "POST"])
|
||||||
def account_group_edit(request, pk):
|
def account_group_edit(request, pk):
|
||||||
account_group = get_object_or_404(AccountGroup, id=pk)
|
account_group = get_shared_object_or_error(AccountGroup, request, id=pk, level=EDIT)
|
||||||
|
|
||||||
if account_group.owner and account_group.owner != request.user:
|
|
||||||
messages.error(request, _("Only the owner can edit this"))
|
|
||||||
|
|
||||||
return HttpResponse(
|
|
||||||
status=204,
|
|
||||||
headers={
|
|
||||||
"HX-Trigger": "updated, hide_offcanvas",
|
|
||||||
},
|
|
||||||
)
|
|
||||||
|
|
||||||
if request.method == "POST":
|
if request.method == "POST":
|
||||||
form = AccountGroupForm(request.POST, instance=account_group)
|
form = AccountGroupForm(request.POST, instance=account_group)
|
||||||
@@ -101,17 +97,18 @@ def account_group_edit(request, pk):
|
|||||||
@login_required
|
@login_required
|
||||||
@require_http_methods(["DELETE"])
|
@require_http_methods(["DELETE"])
|
||||||
def account_group_delete(request, pk):
|
def account_group_delete(request, pk):
|
||||||
account_group = get_object_or_404(AccountGroup, id=pk)
|
account_group = get_shared_object_or_error(AccountGroup, request, id=pk, level=READ)
|
||||||
|
|
||||||
if (
|
if account_group.is_editable_by(request.user):
|
||||||
account_group.owner != request.user
|
account_group.delete()
|
||||||
and request.user in account_group.shared_with.all()
|
messages.success(request, _("Account Group deleted successfully"))
|
||||||
):
|
elif account_group.shared_with.filter(pk=request.user.pk).exists():
|
||||||
|
# Someone else's object shared with us: we can drop our own access
|
||||||
|
# to it, but never delete it.
|
||||||
account_group.shared_with.remove(request.user)
|
account_group.shared_with.remove(request.user)
|
||||||
messages.success(request, _("Item no longer shared with you"))
|
messages.success(request, _("Item no longer shared with you"))
|
||||||
else:
|
else:
|
||||||
account_group.delete()
|
raise PermissionDenied
|
||||||
messages.success(request, _("Account Group deleted successfully"))
|
|
||||||
|
|
||||||
return HttpResponse(
|
return HttpResponse(
|
||||||
status=204,
|
status=204,
|
||||||
@@ -125,7 +122,7 @@ def account_group_delete(request, pk):
|
|||||||
@login_required
|
@login_required
|
||||||
@require_http_methods(["GET"])
|
@require_http_methods(["GET"])
|
||||||
def account_group_take_ownership(request, pk):
|
def account_group_take_ownership(request, pk):
|
||||||
account_group = get_object_or_404(AccountGroup, id=pk)
|
account_group = get_shared_object_or_error(AccountGroup, request, id=pk, level=EDIT)
|
||||||
|
|
||||||
if not account_group.owner:
|
if not account_group.owner:
|
||||||
account_group.owner = request.user
|
account_group.owner = request.user
|
||||||
@@ -146,17 +143,7 @@ def account_group_take_ownership(request, pk):
|
|||||||
@login_required
|
@login_required
|
||||||
@require_http_methods(["GET", "POST"])
|
@require_http_methods(["GET", "POST"])
|
||||||
def account_group_share(request, pk):
|
def account_group_share(request, pk):
|
||||||
obj = get_object_or_404(AccountGroup, id=pk)
|
obj = get_shared_object_or_error(AccountGroup, request, id=pk, level=EDIT)
|
||||||
|
|
||||||
if obj.owner and obj.owner != request.user:
|
|
||||||
messages.error(request, _("Only the owner can edit this"))
|
|
||||||
|
|
||||||
return HttpResponse(
|
|
||||||
status=204,
|
|
||||||
headers={
|
|
||||||
"HX-Trigger": "updated, hide_offcanvas",
|
|
||||||
},
|
|
||||||
)
|
|
||||||
|
|
||||||
if request.method == "POST":
|
if request.method == "POST":
|
||||||
form = SharedObjectForm(request.POST, instance=obj, user=request.user)
|
form = SharedObjectForm(request.POST, instance=obj, user=request.user)
|
||||||
|
|||||||
@@ -1,13 +1,19 @@
|
|||||||
from django.contrib import messages
|
from django.contrib import messages
|
||||||
from django.contrib.auth.decorators import login_required
|
from django.contrib.auth.decorators import login_required
|
||||||
|
from django.core.exceptions import PermissionDenied
|
||||||
from django.http import HttpResponse
|
from django.http import HttpResponse
|
||||||
from django.shortcuts import render, get_object_or_404
|
from django.shortcuts import render
|
||||||
from django.utils.translation import gettext_lazy as _
|
from django.utils.translation import gettext_lazy as _
|
||||||
from django.views.decorators.http import require_http_methods
|
from django.views.decorators.http import require_http_methods
|
||||||
|
|
||||||
from apps.accounts.forms import AccountForm
|
from apps.accounts.forms import AccountForm
|
||||||
from apps.accounts.models import Account
|
from apps.accounts.models import Account
|
||||||
from apps.common.decorators.htmx import only_htmx
|
from apps.common.decorators.htmx import only_htmx
|
||||||
|
from apps.common.functions.permissions import (
|
||||||
|
EDIT,
|
||||||
|
READ,
|
||||||
|
get_shared_object_or_error,
|
||||||
|
)
|
||||||
from apps.common.models import SharedObject
|
from apps.common.models import SharedObject
|
||||||
from apps.common.forms import SharedObjectForm
|
from apps.common.forms import SharedObjectForm
|
||||||
|
|
||||||
@@ -25,7 +31,7 @@ def accounts_index(request):
|
|||||||
@login_required
|
@login_required
|
||||||
@require_http_methods(["GET"])
|
@require_http_methods(["GET"])
|
||||||
def accounts_list(request):
|
def accounts_list(request):
|
||||||
accounts = Account.objects.all().order_by("id")
|
accounts = Account.objects.all().order_by("name")
|
||||||
return render(
|
return render(
|
||||||
request,
|
request,
|
||||||
"accounts/fragments/list.html",
|
"accounts/fragments/list.html",
|
||||||
@@ -63,16 +69,7 @@ def account_add(request, **kwargs):
|
|||||||
@login_required
|
@login_required
|
||||||
@require_http_methods(["GET", "POST"])
|
@require_http_methods(["GET", "POST"])
|
||||||
def account_edit(request, pk):
|
def account_edit(request, pk):
|
||||||
account = get_object_or_404(Account, id=pk)
|
account = get_shared_object_or_error(Account, request, id=pk, level=EDIT)
|
||||||
if account.owner and account.owner != request.user:
|
|
||||||
messages.error(request, _("Only the owner can edit this"))
|
|
||||||
|
|
||||||
return HttpResponse(
|
|
||||||
status=204,
|
|
||||||
headers={
|
|
||||||
"HX-Trigger": "updated, hide_offcanvas",
|
|
||||||
},
|
|
||||||
)
|
|
||||||
|
|
||||||
if request.method == "POST":
|
if request.method == "POST":
|
||||||
form = AccountForm(request.POST, instance=account)
|
form = AccountForm(request.POST, instance=account)
|
||||||
@@ -100,17 +97,7 @@ def account_edit(request, pk):
|
|||||||
@login_required
|
@login_required
|
||||||
@require_http_methods(["GET", "POST"])
|
@require_http_methods(["GET", "POST"])
|
||||||
def account_share(request, pk):
|
def account_share(request, pk):
|
||||||
obj = get_object_or_404(Account, id=pk)
|
obj = get_shared_object_or_error(Account, request, id=pk, level=EDIT)
|
||||||
|
|
||||||
if obj.owner and obj.owner != request.user:
|
|
||||||
messages.error(request, _("Only the owner can edit this"))
|
|
||||||
|
|
||||||
return HttpResponse(
|
|
||||||
status=204,
|
|
||||||
headers={
|
|
||||||
"HX-Trigger": "updated, hide_offcanvas",
|
|
||||||
},
|
|
||||||
)
|
|
||||||
|
|
||||||
if request.method == "POST":
|
if request.method == "POST":
|
||||||
form = SharedObjectForm(request.POST, instance=obj, user=request.user)
|
form = SharedObjectForm(request.POST, instance=obj, user=request.user)
|
||||||
@@ -138,14 +125,18 @@ def account_share(request, pk):
|
|||||||
@login_required
|
@login_required
|
||||||
@require_http_methods(["DELETE"])
|
@require_http_methods(["DELETE"])
|
||||||
def account_delete(request, pk):
|
def account_delete(request, pk):
|
||||||
account = get_object_or_404(Account, id=pk)
|
account = get_shared_object_or_error(Account, request, id=pk, level=READ)
|
||||||
|
|
||||||
if account.owner != request.user and request.user in account.shared_with.all():
|
if account.is_editable_by(request.user):
|
||||||
|
account.delete()
|
||||||
|
messages.success(request, _("Account deleted successfully"))
|
||||||
|
elif account.shared_with.filter(pk=request.user.pk).exists():
|
||||||
|
# Someone else's object shared with us: we can drop our own access
|
||||||
|
# to it, but never delete it.
|
||||||
account.shared_with.remove(request.user)
|
account.shared_with.remove(request.user)
|
||||||
messages.success(request, _("Item no longer shared with you"))
|
messages.success(request, _("Item no longer shared with you"))
|
||||||
else:
|
else:
|
||||||
account.delete()
|
raise PermissionDenied
|
||||||
messages.success(request, _("Account deleted successfully"))
|
|
||||||
|
|
||||||
return HttpResponse(
|
return HttpResponse(
|
||||||
status=204,
|
status=204,
|
||||||
@@ -159,7 +150,9 @@ def account_delete(request, pk):
|
|||||||
@login_required
|
@login_required
|
||||||
@require_http_methods(["GET"])
|
@require_http_methods(["GET"])
|
||||||
def account_toggle_untracked(request, pk):
|
def account_toggle_untracked(request, pk):
|
||||||
account = get_object_or_404(Account, id=pk)
|
# Only flips the calling user's own row in untracked_by, so visibility --
|
||||||
|
# not ownership -- is the right bar here.
|
||||||
|
account = get_shared_object_or_error(Account, request, id=pk, level=READ)
|
||||||
if account.is_untracked_by():
|
if account.is_untracked_by():
|
||||||
account.untracked_by.remove(request.user)
|
account.untracked_by.remove(request.user)
|
||||||
messages.success(request, _("Account is now tracked"))
|
messages.success(request, _("Account is now tracked"))
|
||||||
@@ -179,7 +172,7 @@ def account_toggle_untracked(request, pk):
|
|||||||
@login_required
|
@login_required
|
||||||
@require_http_methods(["GET"])
|
@require_http_methods(["GET"])
|
||||||
def account_take_ownership(request, pk):
|
def account_take_ownership(request, pk):
|
||||||
account = get_object_or_404(Account, id=pk)
|
account = get_shared_object_or_error(Account, request, id=pk, level=EDIT)
|
||||||
|
|
||||||
if not account.owner:
|
if not account.owner:
|
||||||
account.owner = request.user
|
account.owner = request.user
|
||||||
|
|||||||
@@ -0,0 +1,64 @@
|
|||||||
|
from datetime import timedelta
|
||||||
|
|
||||||
|
from django.conf import settings
|
||||||
|
from django.utils import timezone
|
||||||
|
from rest_framework.authentication import BaseAuthentication, get_authorization_header
|
||||||
|
from rest_framework.exceptions import AuthenticationFailed
|
||||||
|
|
||||||
|
from apps.users.models import APIToken
|
||||||
|
|
||||||
|
|
||||||
|
class APITokenAuthentication(BaseAuthentication):
|
||||||
|
keyword = "Token"
|
||||||
|
|
||||||
|
def authenticate(self, request):
|
||||||
|
auth = get_authorization_header(request).split()
|
||||||
|
if not auth or auth[0].lower() != self.keyword.lower().encode():
|
||||||
|
return None
|
||||||
|
|
||||||
|
if len(auth) != 2:
|
||||||
|
raise AuthenticationFailed("Invalid API token header.")
|
||||||
|
|
||||||
|
try:
|
||||||
|
raw_token = auth[1].decode("utf-8")
|
||||||
|
except UnicodeDecodeError as exc:
|
||||||
|
raise AuthenticationFailed("Invalid API token header.") from exc
|
||||||
|
|
||||||
|
# Only claim tokens carrying our prefix; otherwise return None so the
|
||||||
|
# request falls through to other authenticators (e.g. DRF's built-in
|
||||||
|
# TokenAuthentication, which shares the "Token" keyword).
|
||||||
|
if not raw_token.startswith(APIToken.TOKEN_PREFIX):
|
||||||
|
return None
|
||||||
|
|
||||||
|
try:
|
||||||
|
token_key, token_secret = APIToken.parse_raw_token(raw_token)
|
||||||
|
except ValueError as exc:
|
||||||
|
raise AuthenticationFailed("Invalid API token.") from exc
|
||||||
|
|
||||||
|
token = APIToken.objects.select_related("user").filter(token_key=token_key).first()
|
||||||
|
if token is None or not token.check_secret(token_secret):
|
||||||
|
raise AuthenticationFailed("Invalid API token.")
|
||||||
|
if token.revoked_at is not None:
|
||||||
|
raise AuthenticationFailed("API token has been revoked.")
|
||||||
|
if token.is_expired():
|
||||||
|
raise AuthenticationFailed("API token has expired.")
|
||||||
|
if not token.user.is_active:
|
||||||
|
raise AuthenticationFailed("User account is disabled.")
|
||||||
|
|
||||||
|
self._touch_last_used(token)
|
||||||
|
return (token.user, token)
|
||||||
|
|
||||||
|
@staticmethod
|
||||||
|
def _touch_last_used(token):
|
||||||
|
# Avoid a write on every request: only refresh once per interval.
|
||||||
|
now = timezone.now()
|
||||||
|
interval = settings.API_TOKEN_LAST_USED_UPDATE_INTERVAL
|
||||||
|
if (
|
||||||
|
token.last_used_at is None
|
||||||
|
or (now - token.last_used_at) >= timedelta(seconds=interval)
|
||||||
|
):
|
||||||
|
token.last_used_at = now
|
||||||
|
token.save(update_fields=["last_used_at"])
|
||||||
|
|
||||||
|
def authenticate_header(self, request):
|
||||||
|
return self.keyword
|
||||||
@@ -1,4 +1,8 @@
|
|||||||
from rest_framework.permissions import BasePermission
|
from rest_framework.permissions import (
|
||||||
|
SAFE_METHODS,
|
||||||
|
BasePermission,
|
||||||
|
DjangoModelPermissions,
|
||||||
|
)
|
||||||
from django.conf import settings
|
from django.conf import settings
|
||||||
|
|
||||||
|
|
||||||
@@ -8,3 +12,37 @@ class NotInDemoMode(BasePermission):
|
|||||||
return False
|
return False
|
||||||
else:
|
else:
|
||||||
return True
|
return True
|
||||||
|
|
||||||
|
|
||||||
|
class SharedObjectPermission(BasePermission):
|
||||||
|
"""Object-level ownership check for SharedObject-backed viewsets.
|
||||||
|
|
||||||
|
DjangoModelPermissions is model-level: a user holding ``change_account``
|
||||||
|
may write any object the viewset's queryset returns, and for SharedObject
|
||||||
|
that queryset includes other people's public and shared-with-them objects.
|
||||||
|
Sharing grants read access only, so writes are restricted to the owner
|
||||||
|
here as well.
|
||||||
|
|
||||||
|
Set ``shared_object_via`` on the viewset when the governing SharedObject is
|
||||||
|
reached through a relation (e.g. ``"strategy"`` for a DCA entry).
|
||||||
|
"""
|
||||||
|
|
||||||
|
def has_object_permission(self, request, view, obj):
|
||||||
|
if request.method in SAFE_METHODS:
|
||||||
|
return True
|
||||||
|
|
||||||
|
guard = obj
|
||||||
|
via = getattr(view, "shared_object_via", None)
|
||||||
|
for attr in via.split(".") if via else []:
|
||||||
|
guard = getattr(guard, attr)
|
||||||
|
|
||||||
|
return guard.is_editable_by(request.user)
|
||||||
|
|
||||||
|
|
||||||
|
#: Default permissions plus the object-level ownership check. Assigning
|
||||||
|
#: ``permission_classes`` replaces the defaults, so they are repeated here.
|
||||||
|
SHARED_OBJECT_PERMISSIONS = [
|
||||||
|
NotInDemoMode,
|
||||||
|
DjangoModelPermissions,
|
||||||
|
SharedObjectPermission,
|
||||||
|
]
|
||||||
|
|||||||
@@ -1,4 +1,6 @@
|
|||||||
# Import all test classes for Django test discovery
|
# Import all test classes for Django test discovery
|
||||||
from .test_imports import *
|
from .test_imports import *
|
||||||
from .test_accounts import *
|
from .test_accounts import *
|
||||||
|
from .test_data_isolation import *
|
||||||
|
from .test_shared_access import *
|
||||||
|
from .test_object_permissions import *
|
||||||
|
|||||||
@@ -90,10 +90,10 @@ class AccountBalanceAPITests(TestCase):
|
|||||||
self.assertEqual(response.status_code, status.HTTP_404_NOT_FOUND)
|
self.assertEqual(response.status_code, status.HTTP_404_NOT_FOUND)
|
||||||
|
|
||||||
def test_get_balance_unauthenticated(self):
|
def test_get_balance_unauthenticated(self):
|
||||||
"""Test unauthenticated request returns 403"""
|
"""Test unauthenticated request returns 401"""
|
||||||
unauthenticated_client = APIClient()
|
unauthenticated_client = APIClient()
|
||||||
response = unauthenticated_client.get(
|
response = unauthenticated_client.get(
|
||||||
f"/api/accounts/{self.account.id}/balance/"
|
f"/api/accounts/{self.account.id}/balance/"
|
||||||
)
|
)
|
||||||
|
|
||||||
self.assertEqual(response.status_code, status.HTTP_403_FORBIDDEN)
|
self.assertEqual(response.status_code, status.HTTP_401_UNAUTHORIZED)
|
||||||
|
|||||||
@@ -0,0 +1,109 @@
|
|||||||
|
from datetime import timedelta
|
||||||
|
|
||||||
|
from django.contrib.auth import get_user_model
|
||||||
|
from django.test import RequestFactory, TestCase, override_settings
|
||||||
|
from django.utils import timezone
|
||||||
|
from rest_framework.exceptions import AuthenticationFailed
|
||||||
|
|
||||||
|
from apps.api.authentication import APITokenAuthentication
|
||||||
|
from apps.users.models import APIToken
|
||||||
|
|
||||||
|
|
||||||
|
class APITokenAuthenticationTests(TestCase):
|
||||||
|
def setUp(self):
|
||||||
|
self.factory = RequestFactory()
|
||||||
|
self.authentication = APITokenAuthentication()
|
||||||
|
self.user = get_user_model().objects.create_user(
|
||||||
|
email="automation@example.com",
|
||||||
|
password="test-password",
|
||||||
|
)
|
||||||
|
|
||||||
|
def test_returns_none_without_token_header(self):
|
||||||
|
request = self.factory.get("/api/accounts/")
|
||||||
|
self.assertIsNone(self.authentication.authenticate(request))
|
||||||
|
|
||||||
|
def test_authenticates_valid_api_token(self):
|
||||||
|
token, raw_token = APIToken.objects.create_token(user=self.user, name="n8n")
|
||||||
|
request = self.factory.get(
|
||||||
|
"/api/accounts/",
|
||||||
|
HTTP_AUTHORIZATION=f"Token {raw_token}",
|
||||||
|
)
|
||||||
|
|
||||||
|
authenticated_user, authenticated_token = self.authentication.authenticate(request)
|
||||||
|
|
||||||
|
self.assertEqual(authenticated_user, self.user)
|
||||||
|
self.assertEqual(authenticated_token.pk, token.pk)
|
||||||
|
token.refresh_from_db()
|
||||||
|
self.assertIsNotNone(token.last_used_at)
|
||||||
|
|
||||||
|
def test_rejects_expired_api_token(self):
|
||||||
|
token, raw_token = APIToken.objects.create_token(user=self.user, name="n8n")
|
||||||
|
token.expires_at = timezone.now() - timedelta(minutes=1)
|
||||||
|
token.save(update_fields=["expires_at"])
|
||||||
|
request = self.factory.get(
|
||||||
|
"/api/accounts/",
|
||||||
|
HTTP_AUTHORIZATION=f"Token {raw_token}",
|
||||||
|
)
|
||||||
|
|
||||||
|
with self.assertRaisesRegex(AuthenticationFailed, "expired"):
|
||||||
|
self.authentication.authenticate(request)
|
||||||
|
|
||||||
|
def test_rejects_revoked_api_token(self):
|
||||||
|
token, raw_token = APIToken.objects.create_token(user=self.user, name="n8n")
|
||||||
|
token.revoked_at = timezone.now()
|
||||||
|
token.save(update_fields=["revoked_at"])
|
||||||
|
request = self.factory.get(
|
||||||
|
"/api/accounts/",
|
||||||
|
HTTP_AUTHORIZATION=f"Token {raw_token}",
|
||||||
|
)
|
||||||
|
|
||||||
|
with self.assertRaisesRegex(AuthenticationFailed, "revoked"):
|
||||||
|
self.authentication.authenticate(request)
|
||||||
|
|
||||||
|
def test_stores_secret_as_sha256_not_raw(self):
|
||||||
|
token, raw_token = APIToken.objects.create_token(user=self.user, name="n8n")
|
||||||
|
_key, secret = APIToken.parse_raw_token(raw_token)
|
||||||
|
|
||||||
|
self.assertNotIn(secret, token.token_hash)
|
||||||
|
self.assertEqual(len(token.token_hash), 64)
|
||||||
|
self.assertTrue(token.check_secret(secret))
|
||||||
|
|
||||||
|
def test_falls_through_for_non_prefixed_token(self):
|
||||||
|
request = self.factory.get(
|
||||||
|
"/api/accounts/",
|
||||||
|
HTTP_AUTHORIZATION="Token deadbeefdeadbeefdeadbeef",
|
||||||
|
)
|
||||||
|
# Not our prefix: return None so another authenticator can handle it.
|
||||||
|
self.assertIsNone(self.authentication.authenticate(request))
|
||||||
|
|
||||||
|
@override_settings(API_TOKEN_LAST_USED_UPDATE_INTERVAL=600)
|
||||||
|
def test_last_used_at_is_throttled_within_interval(self):
|
||||||
|
token, raw_token = APIToken.objects.create_token(user=self.user, name="n8n")
|
||||||
|
request = self.factory.get(
|
||||||
|
"/api/accounts/",
|
||||||
|
HTTP_AUTHORIZATION=f"Token {raw_token}",
|
||||||
|
)
|
||||||
|
|
||||||
|
self.authentication.authenticate(request)
|
||||||
|
token.refresh_from_db()
|
||||||
|
first_used = token.last_used_at
|
||||||
|
self.assertIsNotNone(first_used)
|
||||||
|
|
||||||
|
self.authentication.authenticate(request)
|
||||||
|
token.refresh_from_db()
|
||||||
|
self.assertEqual(token.last_used_at, first_used)
|
||||||
|
|
||||||
|
@override_settings(API_TOKEN_LAST_USED_UPDATE_INTERVAL=0)
|
||||||
|
def test_last_used_at_updates_after_interval(self):
|
||||||
|
token, raw_token = APIToken.objects.create_token(user=self.user, name="n8n")
|
||||||
|
token.last_used_at = timezone.now() - timedelta(minutes=5)
|
||||||
|
token.save(update_fields=["last_used_at"])
|
||||||
|
stale = token.last_used_at
|
||||||
|
request = self.factory.get(
|
||||||
|
"/api/accounts/",
|
||||||
|
HTTP_AUTHORIZATION=f"Token {raw_token}",
|
||||||
|
)
|
||||||
|
|
||||||
|
self.authentication.authenticate(request)
|
||||||
|
token.refresh_from_db()
|
||||||
|
self.assertGreater(token.last_used_at, stale)
|
||||||
@@ -0,0 +1,719 @@
|
|||||||
|
from datetime import date
|
||||||
|
from decimal import Decimal
|
||||||
|
|
||||||
|
from django.contrib.auth import get_user_model
|
||||||
|
from django.test import TestCase, override_settings
|
||||||
|
from rest_framework import status
|
||||||
|
from rest_framework.test import APIClient
|
||||||
|
|
||||||
|
from apps.accounts.models import Account, AccountGroup
|
||||||
|
from apps.currencies.models import Currency
|
||||||
|
from apps.dca.models import DCAStrategy, DCAEntry
|
||||||
|
from apps.transactions.models import (
|
||||||
|
Transaction,
|
||||||
|
TransactionCategory,
|
||||||
|
TransactionTag,
|
||||||
|
TransactionEntity,
|
||||||
|
InstallmentPlan,
|
||||||
|
RecurringTransaction,
|
||||||
|
)
|
||||||
|
|
||||||
|
|
||||||
|
ACCESS_DENIED_CODES = [status.HTTP_403_FORBIDDEN, status.HTTP_404_NOT_FOUND]
|
||||||
|
|
||||||
|
|
||||||
|
@override_settings(
|
||||||
|
STORAGES={
|
||||||
|
"default": {"BACKEND": "django.core.files.storage.FileSystemStorage"},
|
||||||
|
"staticfiles": {
|
||||||
|
"BACKEND": "django.contrib.staticfiles.storage.StaticFilesStorage"
|
||||||
|
},
|
||||||
|
},
|
||||||
|
WHITENOISE_AUTOREFRESH=True,
|
||||||
|
)
|
||||||
|
class AccountDataIsolationTests(TestCase):
|
||||||
|
"""Tests to ensure users cannot access other users' accounts."""
|
||||||
|
|
||||||
|
def setUp(self):
|
||||||
|
"""Set up test data with two distinct users."""
|
||||||
|
User = get_user_model()
|
||||||
|
|
||||||
|
# User 1 - the requester
|
||||||
|
self.user1 = User.objects.create_user(
|
||||||
|
email="user1@test.com", password="testpass123"
|
||||||
|
)
|
||||||
|
self.client1 = APIClient()
|
||||||
|
self.client1.force_authenticate(user=self.user1)
|
||||||
|
|
||||||
|
# User 2 - owner of data that user1 should NOT access
|
||||||
|
self.user2 = User.objects.create_user(
|
||||||
|
email="user2@test.com", password="testpass123"
|
||||||
|
)
|
||||||
|
self.client2 = APIClient()
|
||||||
|
self.client2.force_authenticate(user=self.user2)
|
||||||
|
|
||||||
|
# Shared currency
|
||||||
|
self.currency = Currency.objects.create(
|
||||||
|
code="USD", name="US Dollar", decimal_places=2, prefix="$ "
|
||||||
|
)
|
||||||
|
|
||||||
|
# User 1's account
|
||||||
|
self.user1_account_group = AccountGroup.all_objects.create(
|
||||||
|
name="User1 Group", owner=self.user1
|
||||||
|
)
|
||||||
|
self.user1_account = Account.all_objects.create(
|
||||||
|
name="User1 Account",
|
||||||
|
group=self.user1_account_group,
|
||||||
|
currency=self.currency,
|
||||||
|
owner=self.user1,
|
||||||
|
)
|
||||||
|
|
||||||
|
# User 2's account (private, should be invisible to user1)
|
||||||
|
self.user2_account_group = AccountGroup.all_objects.create(
|
||||||
|
name="User2 Group", owner=self.user2
|
||||||
|
)
|
||||||
|
self.user2_account = Account.all_objects.create(
|
||||||
|
name="User2 Account",
|
||||||
|
group=self.user2_account_group,
|
||||||
|
currency=self.currency,
|
||||||
|
owner=self.user2,
|
||||||
|
)
|
||||||
|
|
||||||
|
def test_user_cannot_see_other_users_accounts_in_list(self):
|
||||||
|
"""GET /api/accounts/ should only return user's own accounts."""
|
||||||
|
response = self.client1.get("/api/accounts/")
|
||||||
|
|
||||||
|
self.assertEqual(response.status_code, status.HTTP_200_OK)
|
||||||
|
|
||||||
|
# User1 should only see their own account
|
||||||
|
account_ids = [acc["id"] for acc in response.data["results"]]
|
||||||
|
self.assertIn(self.user1_account.id, account_ids)
|
||||||
|
self.assertNotIn(self.user2_account.id, account_ids)
|
||||||
|
|
||||||
|
def test_user_cannot_access_other_users_account_detail(self):
|
||||||
|
"""GET /api/accounts/{id}/ should deny access to other user's account."""
|
||||||
|
response = self.client1.get(f"/api/accounts/{self.user2_account.id}/")
|
||||||
|
|
||||||
|
self.assertIn(response.status_code, ACCESS_DENIED_CODES)
|
||||||
|
|
||||||
|
def test_user_cannot_modify_other_users_account(self):
|
||||||
|
"""PATCH on other user's account should deny access."""
|
||||||
|
response = self.client1.patch(
|
||||||
|
f"/api/accounts/{self.user2_account.id}/",
|
||||||
|
{"name": "Hacked Account"},
|
||||||
|
)
|
||||||
|
self.assertIn(response.status_code, ACCESS_DENIED_CODES)
|
||||||
|
|
||||||
|
# Verify account name wasn't changed
|
||||||
|
self.user2_account.refresh_from_db()
|
||||||
|
self.assertEqual(self.user2_account.name, "User2 Account")
|
||||||
|
|
||||||
|
def test_user_cannot_delete_other_users_account(self):
|
||||||
|
"""DELETE on other user's account should deny access."""
|
||||||
|
response = self.client1.delete(f"/api/accounts/{self.user2_account.id}/")
|
||||||
|
|
||||||
|
self.assertIn(response.status_code, ACCESS_DENIED_CODES)
|
||||||
|
|
||||||
|
# Verify account still exists
|
||||||
|
self.assertTrue(Account.all_objects.filter(id=self.user2_account.id).exists())
|
||||||
|
|
||||||
|
def test_user_cannot_get_balance_of_other_users_account(self):
|
||||||
|
"""Balance action on other user's account should deny access."""
|
||||||
|
response = self.client1.get(f"/api/accounts/{self.user2_account.id}/balance/")
|
||||||
|
|
||||||
|
self.assertIn(response.status_code, ACCESS_DENIED_CODES)
|
||||||
|
|
||||||
|
def test_user_can_access_own_account(self):
|
||||||
|
"""User can access their own account normally."""
|
||||||
|
response = self.client1.get(f"/api/accounts/{self.user1_account.id}/")
|
||||||
|
|
||||||
|
self.assertEqual(response.status_code, status.HTTP_200_OK)
|
||||||
|
self.assertEqual(response.data["name"], "User1 Account")
|
||||||
|
|
||||||
|
|
||||||
|
@override_settings(
|
||||||
|
STORAGES={
|
||||||
|
"default": {"BACKEND": "django.core.files.storage.FileSystemStorage"},
|
||||||
|
"staticfiles": {
|
||||||
|
"BACKEND": "django.contrib.staticfiles.storage.StaticFilesStorage"
|
||||||
|
},
|
||||||
|
},
|
||||||
|
WHITENOISE_AUTOREFRESH=True,
|
||||||
|
)
|
||||||
|
class AccountGroupDataIsolationTests(TestCase):
|
||||||
|
"""Tests to ensure users cannot access other users' account groups."""
|
||||||
|
|
||||||
|
def setUp(self):
|
||||||
|
"""Set up test data with two distinct users."""
|
||||||
|
User = get_user_model()
|
||||||
|
|
||||||
|
self.user1 = User.objects.create_user(
|
||||||
|
email="user1@test.com", password="testpass123"
|
||||||
|
)
|
||||||
|
self.client1 = APIClient()
|
||||||
|
self.client1.force_authenticate(user=self.user1)
|
||||||
|
|
||||||
|
self.user2 = User.objects.create_user(
|
||||||
|
email="user2@test.com", password="testpass123"
|
||||||
|
)
|
||||||
|
|
||||||
|
# User 1's account group
|
||||||
|
self.user1_group = AccountGroup.all_objects.create(
|
||||||
|
name="User1 Group", owner=self.user1
|
||||||
|
)
|
||||||
|
|
||||||
|
# User 2's account group
|
||||||
|
self.user2_group = AccountGroup.all_objects.create(
|
||||||
|
name="User2 Group", owner=self.user2
|
||||||
|
)
|
||||||
|
|
||||||
|
def test_user_cannot_see_other_users_account_groups(self):
|
||||||
|
"""GET /api/account-groups/ should only return user's own groups."""
|
||||||
|
response = self.client1.get("/api/account-groups/")
|
||||||
|
|
||||||
|
self.assertEqual(response.status_code, status.HTTP_200_OK)
|
||||||
|
|
||||||
|
group_ids = [grp["id"] for grp in response.data["results"]]
|
||||||
|
self.assertIn(self.user1_group.id, group_ids)
|
||||||
|
self.assertNotIn(self.user2_group.id, group_ids)
|
||||||
|
|
||||||
|
def test_user_cannot_access_other_users_account_group_detail(self):
|
||||||
|
"""GET /api/account-groups/{id}/ should deny access to other user's group."""
|
||||||
|
response = self.client1.get(f"/api/account-groups/{self.user2_group.id}/")
|
||||||
|
|
||||||
|
self.assertIn(response.status_code, ACCESS_DENIED_CODES)
|
||||||
|
|
||||||
|
def test_user_cannot_modify_other_users_account_group(self):
|
||||||
|
"""PATCH on other user's account group should deny access."""
|
||||||
|
response = self.client1.patch(
|
||||||
|
f"/api/account-groups/{self.user2_group.id}/",
|
||||||
|
{"name": "Hacked Group"},
|
||||||
|
)
|
||||||
|
|
||||||
|
self.assertIn(response.status_code, ACCESS_DENIED_CODES)
|
||||||
|
|
||||||
|
self.user2_group.refresh_from_db()
|
||||||
|
self.assertEqual(self.user2_group.name, "User2 Group")
|
||||||
|
|
||||||
|
def test_user_cannot_delete_other_users_account_group(self):
|
||||||
|
"""DELETE on other user's account group should deny access."""
|
||||||
|
response = self.client1.delete(f"/api/account-groups/{self.user2_group.id}/")
|
||||||
|
|
||||||
|
self.assertIn(response.status_code, ACCESS_DENIED_CODES)
|
||||||
|
|
||||||
|
self.assertTrue(
|
||||||
|
AccountGroup.all_objects.filter(id=self.user2_group.id).exists()
|
||||||
|
)
|
||||||
|
|
||||||
|
|
||||||
|
@override_settings(
|
||||||
|
STORAGES={
|
||||||
|
"default": {"BACKEND": "django.core.files.storage.FileSystemStorage"},
|
||||||
|
"staticfiles": {
|
||||||
|
"BACKEND": "django.contrib.staticfiles.storage.StaticFilesStorage"
|
||||||
|
},
|
||||||
|
},
|
||||||
|
WHITENOISE_AUTOREFRESH=True,
|
||||||
|
)
|
||||||
|
class TransactionDataIsolationTests(TestCase):
|
||||||
|
"""Tests to ensure users cannot access other users' transactions."""
|
||||||
|
|
||||||
|
def setUp(self):
|
||||||
|
"""Set up test data with transactions for two distinct users."""
|
||||||
|
User = get_user_model()
|
||||||
|
|
||||||
|
self.user1 = User.objects.create_user(
|
||||||
|
email="user1@test.com", password="testpass123"
|
||||||
|
)
|
||||||
|
self.client1 = APIClient()
|
||||||
|
self.client1.force_authenticate(user=self.user1)
|
||||||
|
|
||||||
|
self.user2 = User.objects.create_user(
|
||||||
|
email="user2@test.com", password="testpass123"
|
||||||
|
)
|
||||||
|
|
||||||
|
self.currency = Currency.objects.create(
|
||||||
|
code="USD", name="US Dollar", decimal_places=2, prefix="$ "
|
||||||
|
)
|
||||||
|
|
||||||
|
# User 1's account and transaction
|
||||||
|
self.user1_account = Account.all_objects.create(
|
||||||
|
name="User1 Account", currency=self.currency, owner=self.user1
|
||||||
|
)
|
||||||
|
self.user1_transaction = Transaction.userless_all_objects.create(
|
||||||
|
account=self.user1_account,
|
||||||
|
type=Transaction.Type.INCOME,
|
||||||
|
amount=Decimal("100.00"),
|
||||||
|
is_paid=True,
|
||||||
|
date=date(2025, 1, 1),
|
||||||
|
description="User1 Income",
|
||||||
|
owner=self.user1,
|
||||||
|
)
|
||||||
|
|
||||||
|
# User 2's account and transaction
|
||||||
|
self.user2_account = Account.all_objects.create(
|
||||||
|
name="User2 Account", currency=self.currency, owner=self.user2
|
||||||
|
)
|
||||||
|
self.user2_transaction = Transaction.userless_all_objects.create(
|
||||||
|
account=self.user2_account,
|
||||||
|
type=Transaction.Type.EXPENSE,
|
||||||
|
amount=Decimal("50.00"),
|
||||||
|
is_paid=True,
|
||||||
|
date=date(2025, 1, 1),
|
||||||
|
description="User2 Expense",
|
||||||
|
owner=self.user2,
|
||||||
|
)
|
||||||
|
|
||||||
|
def test_user_cannot_see_other_users_transactions_in_list(self):
|
||||||
|
"""GET /api/transactions/ should only return user's own transactions."""
|
||||||
|
response = self.client1.get("/api/transactions/")
|
||||||
|
|
||||||
|
self.assertEqual(response.status_code, status.HTTP_200_OK)
|
||||||
|
|
||||||
|
transaction_ids = [t["id"] for t in response.data["results"]]
|
||||||
|
self.assertIn(self.user1_transaction.id, transaction_ids)
|
||||||
|
self.assertNotIn(self.user2_transaction.id, transaction_ids)
|
||||||
|
|
||||||
|
def test_user_cannot_access_other_users_transaction_detail(self):
|
||||||
|
"""GET /api/transactions/{id}/ should deny access to other user's transaction."""
|
||||||
|
response = self.client1.get(f"/api/transactions/{self.user2_transaction.id}/")
|
||||||
|
|
||||||
|
self.assertIn(response.status_code, ACCESS_DENIED_CODES)
|
||||||
|
|
||||||
|
def test_user_cannot_modify_other_users_transaction(self):
|
||||||
|
"""PATCH on other user's transaction should deny access."""
|
||||||
|
response = self.client1.patch(
|
||||||
|
f"/api/transactions/{self.user2_transaction.id}/",
|
||||||
|
{"description": "Hacked Transaction"},
|
||||||
|
)
|
||||||
|
|
||||||
|
self.assertIn(response.status_code, ACCESS_DENIED_CODES)
|
||||||
|
|
||||||
|
self.user2_transaction.refresh_from_db()
|
||||||
|
self.assertEqual(self.user2_transaction.description, "User2 Expense")
|
||||||
|
|
||||||
|
def test_user_cannot_delete_other_users_transaction(self):
|
||||||
|
"""DELETE on other user's transaction should deny access."""
|
||||||
|
response = self.client1.delete(
|
||||||
|
f"/api/transactions/{self.user2_transaction.id}/"
|
||||||
|
)
|
||||||
|
|
||||||
|
self.assertIn(response.status_code, ACCESS_DENIED_CODES)
|
||||||
|
|
||||||
|
self.assertTrue(
|
||||||
|
Transaction.userless_all_objects.filter(
|
||||||
|
id=self.user2_transaction.id
|
||||||
|
).exists()
|
||||||
|
)
|
||||||
|
|
||||||
|
def test_user_cannot_create_transaction_in_other_users_account(self):
|
||||||
|
"""POST /api/transactions/ with other user's account should fail."""
|
||||||
|
response = self.client1.post(
|
||||||
|
"/api/transactions/",
|
||||||
|
{
|
||||||
|
"account": self.user2_account.id,
|
||||||
|
"type": "IN",
|
||||||
|
"amount": "100.00",
|
||||||
|
"date": "2025-01-15",
|
||||||
|
"description": "Sneaky transaction",
|
||||||
|
},
|
||||||
|
format="json",
|
||||||
|
)
|
||||||
|
|
||||||
|
# Should deny access - 400 (validation error), 403, or 404
|
||||||
|
self.assertIn(
|
||||||
|
response.status_code,
|
||||||
|
ACCESS_DENIED_CODES + [status.HTTP_400_BAD_REQUEST],
|
||||||
|
)
|
||||||
|
|
||||||
|
|
||||||
|
@override_settings(
|
||||||
|
STORAGES={
|
||||||
|
"default": {"BACKEND": "django.core.files.storage.FileSystemStorage"},
|
||||||
|
"staticfiles": {
|
||||||
|
"BACKEND": "django.contrib.staticfiles.storage.StaticFilesStorage"
|
||||||
|
},
|
||||||
|
},
|
||||||
|
WHITENOISE_AUTOREFRESH=True,
|
||||||
|
)
|
||||||
|
class CategoryTagEntityIsolationTests(TestCase):
|
||||||
|
"""Tests for isolation of categories, tags, and entities between users."""
|
||||||
|
|
||||||
|
def setUp(self):
|
||||||
|
"""Set up test data."""
|
||||||
|
User = get_user_model()
|
||||||
|
|
||||||
|
self.user1 = User.objects.create_user(
|
||||||
|
email="user1@test.com", password="testpass123"
|
||||||
|
)
|
||||||
|
self.client1 = APIClient()
|
||||||
|
self.client1.force_authenticate(user=self.user1)
|
||||||
|
|
||||||
|
self.user2 = User.objects.create_user(
|
||||||
|
email="user2@test.com", password="testpass123"
|
||||||
|
)
|
||||||
|
|
||||||
|
# User 1's categories, tags, entities
|
||||||
|
self.user1_category = TransactionCategory.all_objects.create(
|
||||||
|
name="User1 Category", owner=self.user1
|
||||||
|
)
|
||||||
|
self.user1_tag = TransactionTag.all_objects.create(
|
||||||
|
name="User1 Tag", owner=self.user1
|
||||||
|
)
|
||||||
|
self.user1_entity = TransactionEntity.all_objects.create(
|
||||||
|
name="User1 Entity", owner=self.user1
|
||||||
|
)
|
||||||
|
|
||||||
|
# User 2's categories, tags, entities
|
||||||
|
self.user2_category = TransactionCategory.all_objects.create(
|
||||||
|
name="User2 Category", owner=self.user2
|
||||||
|
)
|
||||||
|
self.user2_tag = TransactionTag.all_objects.create(
|
||||||
|
name="User2 Tag", owner=self.user2
|
||||||
|
)
|
||||||
|
self.user2_entity = TransactionEntity.all_objects.create(
|
||||||
|
name="User2 Entity", owner=self.user2
|
||||||
|
)
|
||||||
|
|
||||||
|
def test_user_cannot_see_other_users_categories(self):
|
||||||
|
"""GET /api/categories/ should only return user's own categories."""
|
||||||
|
response = self.client1.get("/api/categories/")
|
||||||
|
|
||||||
|
self.assertEqual(response.status_code, status.HTTP_200_OK)
|
||||||
|
|
||||||
|
category_ids = [c["id"] for c in response.data["results"]]
|
||||||
|
self.assertIn(self.user1_category.id, category_ids)
|
||||||
|
self.assertNotIn(self.user2_category.id, category_ids)
|
||||||
|
|
||||||
|
def test_user_cannot_access_other_users_category_detail(self):
|
||||||
|
"""GET /api/categories/{id}/ should deny access to other user's category."""
|
||||||
|
response = self.client1.get(f"/api/categories/{self.user2_category.id}/")
|
||||||
|
|
||||||
|
self.assertIn(response.status_code, ACCESS_DENIED_CODES)
|
||||||
|
|
||||||
|
def test_user_cannot_see_other_users_tags(self):
|
||||||
|
"""GET /api/tags/ should only return user's own tags."""
|
||||||
|
response = self.client1.get("/api/tags/")
|
||||||
|
|
||||||
|
self.assertEqual(response.status_code, status.HTTP_200_OK)
|
||||||
|
|
||||||
|
tag_ids = [t["id"] for t in response.data["results"]]
|
||||||
|
self.assertIn(self.user1_tag.id, tag_ids)
|
||||||
|
self.assertNotIn(self.user2_tag.id, tag_ids)
|
||||||
|
|
||||||
|
def test_user_cannot_access_other_users_tag_detail(self):
|
||||||
|
"""GET /api/tags/{id}/ should deny access to other user's tag."""
|
||||||
|
response = self.client1.get(f"/api/tags/{self.user2_tag.id}/")
|
||||||
|
|
||||||
|
self.assertIn(response.status_code, ACCESS_DENIED_CODES)
|
||||||
|
|
||||||
|
def test_user_cannot_see_other_users_entities(self):
|
||||||
|
"""GET /api/entities/ should only return user's own entities."""
|
||||||
|
response = self.client1.get("/api/entities/")
|
||||||
|
|
||||||
|
self.assertEqual(response.status_code, status.HTTP_200_OK)
|
||||||
|
|
||||||
|
entity_ids = [e["id"] for e in response.data["results"]]
|
||||||
|
self.assertIn(self.user1_entity.id, entity_ids)
|
||||||
|
self.assertNotIn(self.user2_entity.id, entity_ids)
|
||||||
|
|
||||||
|
def test_user_cannot_access_other_users_entity_detail(self):
|
||||||
|
"""GET /api/entities/{id}/ should deny access to other user's entity."""
|
||||||
|
response = self.client1.get(f"/api/entities/{self.user2_entity.id}/")
|
||||||
|
|
||||||
|
self.assertIn(response.status_code, ACCESS_DENIED_CODES)
|
||||||
|
|
||||||
|
def test_user_cannot_modify_other_users_category(self):
|
||||||
|
"""PATCH on other user's category should deny access."""
|
||||||
|
response = self.client1.patch(
|
||||||
|
f"/api/categories/{self.user2_category.id}/",
|
||||||
|
{"name": "Hacked Category"},
|
||||||
|
)
|
||||||
|
|
||||||
|
self.assertIn(response.status_code, ACCESS_DENIED_CODES)
|
||||||
|
|
||||||
|
def test_user_cannot_delete_other_users_tag(self):
|
||||||
|
"""DELETE on other user's tag should deny access."""
|
||||||
|
response = self.client1.delete(f"/api/tags/{self.user2_tag.id}/")
|
||||||
|
|
||||||
|
self.assertIn(response.status_code, ACCESS_DENIED_CODES)
|
||||||
|
|
||||||
|
self.assertTrue(
|
||||||
|
TransactionTag.all_objects.filter(id=self.user2_tag.id).exists()
|
||||||
|
)
|
||||||
|
|
||||||
|
|
||||||
|
@override_settings(
|
||||||
|
STORAGES={
|
||||||
|
"default": {"BACKEND": "django.core.files.storage.FileSystemStorage"},
|
||||||
|
"staticfiles": {
|
||||||
|
"BACKEND": "django.contrib.staticfiles.storage.StaticFilesStorage"
|
||||||
|
},
|
||||||
|
},
|
||||||
|
WHITENOISE_AUTOREFRESH=True,
|
||||||
|
)
|
||||||
|
class DCADataIsolationTests(TestCase):
|
||||||
|
"""Tests to ensure users cannot access other users' DCA strategies and entries."""
|
||||||
|
|
||||||
|
def setUp(self):
|
||||||
|
"""Set up test data."""
|
||||||
|
User = get_user_model()
|
||||||
|
|
||||||
|
self.user1 = User.objects.create_user(
|
||||||
|
email="user1@test.com", password="testpass123"
|
||||||
|
)
|
||||||
|
self.client1 = APIClient()
|
||||||
|
self.client1.force_authenticate(user=self.user1)
|
||||||
|
|
||||||
|
self.user2 = User.objects.create_user(
|
||||||
|
email="user2@test.com", password="testpass123"
|
||||||
|
)
|
||||||
|
|
||||||
|
self.currency1 = Currency.objects.create(
|
||||||
|
code="BTC", name="Bitcoin", decimal_places=8, prefix=""
|
||||||
|
)
|
||||||
|
self.currency2 = Currency.objects.create(
|
||||||
|
code="USD", name="US Dollar", decimal_places=2, prefix="$ "
|
||||||
|
)
|
||||||
|
|
||||||
|
# User 1's DCA strategy and entry
|
||||||
|
self.user1_strategy = DCAStrategy.all_objects.create(
|
||||||
|
name="User1 BTC Strategy",
|
||||||
|
target_currency=self.currency1,
|
||||||
|
payment_currency=self.currency2,
|
||||||
|
owner=self.user1,
|
||||||
|
)
|
||||||
|
self.user1_entry = DCAEntry.objects.create(
|
||||||
|
strategy=self.user1_strategy,
|
||||||
|
date=date(2025, 1, 1),
|
||||||
|
amount_paid=Decimal("100.00"),
|
||||||
|
amount_received=Decimal("0.001"),
|
||||||
|
)
|
||||||
|
|
||||||
|
# User 2's DCA strategy and entry
|
||||||
|
self.user2_strategy = DCAStrategy.all_objects.create(
|
||||||
|
name="User2 BTC Strategy",
|
||||||
|
target_currency=self.currency1,
|
||||||
|
payment_currency=self.currency2,
|
||||||
|
owner=self.user2,
|
||||||
|
)
|
||||||
|
self.user2_entry = DCAEntry.objects.create(
|
||||||
|
strategy=self.user2_strategy,
|
||||||
|
date=date(2025, 1, 1),
|
||||||
|
amount_paid=Decimal("200.00"),
|
||||||
|
amount_received=Decimal("0.002"),
|
||||||
|
)
|
||||||
|
|
||||||
|
def test_user_cannot_see_other_users_dca_strategies(self):
|
||||||
|
"""GET /api/dca/strategies/ should only return user's own strategies."""
|
||||||
|
response = self.client1.get("/api/dca/strategies/")
|
||||||
|
|
||||||
|
self.assertEqual(response.status_code, status.HTTP_200_OK)
|
||||||
|
|
||||||
|
strategy_ids = [s["id"] for s in response.data["results"]]
|
||||||
|
self.assertIn(self.user1_strategy.id, strategy_ids)
|
||||||
|
self.assertNotIn(self.user2_strategy.id, strategy_ids)
|
||||||
|
|
||||||
|
def test_user_cannot_access_other_users_dca_strategy_detail(self):
|
||||||
|
"""GET /api/dca/strategies/{id}/ should deny access to other user's strategy."""
|
||||||
|
response = self.client1.get(f"/api/dca/strategies/{self.user2_strategy.id}/")
|
||||||
|
|
||||||
|
self.assertIn(response.status_code, ACCESS_DENIED_CODES)
|
||||||
|
|
||||||
|
def test_user_cannot_access_other_users_dca_entries(self):
|
||||||
|
"""GET /api/dca/entries/ filtered by other user's strategy should return empty."""
|
||||||
|
response = self.client1.get(
|
||||||
|
f"/api/dca/entries/?strategy={self.user2_strategy.id}"
|
||||||
|
)
|
||||||
|
|
||||||
|
# Either OK with empty results or error
|
||||||
|
if response.status_code == status.HTTP_200_OK:
|
||||||
|
entry_ids = [e["id"] for e in response.data["results"]]
|
||||||
|
self.assertNotIn(self.user2_entry.id, entry_ids)
|
||||||
|
|
||||||
|
def test_user_cannot_access_other_users_dca_entry_detail(self):
|
||||||
|
"""GET /api/dca/entries/{id}/ should deny access to other user's entry."""
|
||||||
|
response = self.client1.get(f"/api/dca/entries/{self.user2_entry.id}/")
|
||||||
|
|
||||||
|
self.assertIn(response.status_code, ACCESS_DENIED_CODES)
|
||||||
|
|
||||||
|
def test_user_cannot_access_other_users_strategy_investment_frequency(self):
|
||||||
|
"""investment_frequency action on other user's strategy should deny access."""
|
||||||
|
response = self.client1.get(
|
||||||
|
f"/api/dca/strategies/{self.user2_strategy.id}/investment_frequency/"
|
||||||
|
)
|
||||||
|
|
||||||
|
self.assertIn(response.status_code, ACCESS_DENIED_CODES)
|
||||||
|
|
||||||
|
def test_user_cannot_access_other_users_strategy_price_comparison(self):
|
||||||
|
"""price_comparison action on other user's strategy should deny access."""
|
||||||
|
response = self.client1.get(
|
||||||
|
f"/api/dca/strategies/{self.user2_strategy.id}/price_comparison/"
|
||||||
|
)
|
||||||
|
|
||||||
|
self.assertIn(response.status_code, ACCESS_DENIED_CODES)
|
||||||
|
|
||||||
|
def test_user_cannot_access_other_users_strategy_current_price(self):
|
||||||
|
"""current_price action on other user's strategy should deny access."""
|
||||||
|
response = self.client1.get(
|
||||||
|
f"/api/dca/strategies/{self.user2_strategy.id}/current_price/"
|
||||||
|
)
|
||||||
|
|
||||||
|
self.assertIn(response.status_code, ACCESS_DENIED_CODES)
|
||||||
|
|
||||||
|
def test_user_cannot_modify_other_users_dca_strategy(self):
|
||||||
|
"""PATCH on other user's DCA strategy should deny access."""
|
||||||
|
response = self.client1.patch(
|
||||||
|
f"/api/dca/strategies/{self.user2_strategy.id}/",
|
||||||
|
{"name": "Hacked Strategy"},
|
||||||
|
)
|
||||||
|
|
||||||
|
self.assertIn(response.status_code, ACCESS_DENIED_CODES)
|
||||||
|
|
||||||
|
def test_user_cannot_delete_other_users_dca_entry(self):
|
||||||
|
"""DELETE on other user's DCA entry should deny access."""
|
||||||
|
response = self.client1.delete(f"/api/dca/entries/{self.user2_entry.id}/")
|
||||||
|
|
||||||
|
self.assertIn(response.status_code, ACCESS_DENIED_CODES)
|
||||||
|
|
||||||
|
self.assertTrue(DCAEntry.objects.filter(id=self.user2_entry.id).exists())
|
||||||
|
|
||||||
|
|
||||||
|
@override_settings(
|
||||||
|
STORAGES={
|
||||||
|
"default": {"BACKEND": "django.core.files.storage.FileSystemStorage"},
|
||||||
|
"staticfiles": {
|
||||||
|
"BACKEND": "django.contrib.staticfiles.storage.StaticFilesStorage"
|
||||||
|
},
|
||||||
|
},
|
||||||
|
WHITENOISE_AUTOREFRESH=True,
|
||||||
|
)
|
||||||
|
class InstallmentRecurringIsolationTests(TestCase):
|
||||||
|
"""Tests for isolation of installment plans and recurring transactions."""
|
||||||
|
|
||||||
|
def setUp(self):
|
||||||
|
"""Set up test data."""
|
||||||
|
User = get_user_model()
|
||||||
|
|
||||||
|
self.user1 = User.objects.create_user(
|
||||||
|
email="user1@test.com", password="testpass123"
|
||||||
|
)
|
||||||
|
self.client1 = APIClient()
|
||||||
|
self.client1.force_authenticate(user=self.user1)
|
||||||
|
|
||||||
|
self.user2 = User.objects.create_user(
|
||||||
|
email="user2@test.com", password="testpass123"
|
||||||
|
)
|
||||||
|
|
||||||
|
self.currency = Currency.objects.create(
|
||||||
|
code="USD", name="US Dollar", decimal_places=2, prefix="$ "
|
||||||
|
)
|
||||||
|
|
||||||
|
# User 1's account
|
||||||
|
self.user1_account = Account.all_objects.create(
|
||||||
|
name="User1 Account", currency=self.currency, owner=self.user1
|
||||||
|
)
|
||||||
|
|
||||||
|
# User 2's account
|
||||||
|
self.user2_account = Account.all_objects.create(
|
||||||
|
name="User2 Account", currency=self.currency, owner=self.user2
|
||||||
|
)
|
||||||
|
|
||||||
|
# User 1's installment plan
|
||||||
|
self.user1_installment = InstallmentPlan.all_objects.create(
|
||||||
|
account=self.user1_account,
|
||||||
|
type=Transaction.Type.EXPENSE,
|
||||||
|
description="User1 Installment",
|
||||||
|
number_of_installments=12,
|
||||||
|
start_date=date(2025, 1, 1),
|
||||||
|
installment_amount=Decimal("100.00"),
|
||||||
|
)
|
||||||
|
|
||||||
|
# User 2's installment plan
|
||||||
|
self.user2_installment = InstallmentPlan.all_objects.create(
|
||||||
|
account=self.user2_account,
|
||||||
|
type=Transaction.Type.EXPENSE,
|
||||||
|
description="User2 Installment",
|
||||||
|
number_of_installments=6,
|
||||||
|
start_date=date(2025, 1, 1),
|
||||||
|
installment_amount=Decimal("200.00"),
|
||||||
|
)
|
||||||
|
|
||||||
|
# User 1's recurring transaction
|
||||||
|
self.user1_recurring = RecurringTransaction.all_objects.create(
|
||||||
|
account=self.user1_account,
|
||||||
|
type=Transaction.Type.EXPENSE,
|
||||||
|
amount=Decimal("50.00"),
|
||||||
|
description="User1 Recurring",
|
||||||
|
start_date=date(2025, 1, 1),
|
||||||
|
recurrence_type=RecurringTransaction.RecurrenceType.MONTH,
|
||||||
|
recurrence_interval=1,
|
||||||
|
)
|
||||||
|
|
||||||
|
# User 2's recurring transaction
|
||||||
|
self.user2_recurring = RecurringTransaction.all_objects.create(
|
||||||
|
account=self.user2_account,
|
||||||
|
type=Transaction.Type.INCOME,
|
||||||
|
amount=Decimal("1000.00"),
|
||||||
|
description="User2 Recurring",
|
||||||
|
start_date=date(2025, 1, 1),
|
||||||
|
recurrence_type=RecurringTransaction.RecurrenceType.MONTH,
|
||||||
|
recurrence_interval=1,
|
||||||
|
)
|
||||||
|
|
||||||
|
def test_user_cannot_see_other_users_installment_plans(self):
|
||||||
|
"""GET /api/installment-plans/ should only return user's own plans."""
|
||||||
|
response = self.client1.get("/api/installment-plans/")
|
||||||
|
|
||||||
|
self.assertEqual(response.status_code, status.HTTP_200_OK)
|
||||||
|
|
||||||
|
plan_ids = [p["id"] for p in response.data["results"]]
|
||||||
|
self.assertIn(self.user1_installment.id, plan_ids)
|
||||||
|
self.assertNotIn(self.user2_installment.id, plan_ids)
|
||||||
|
|
||||||
|
def test_user_cannot_access_other_users_installment_plan_detail(self):
|
||||||
|
"""GET /api/installment-plans/{id}/ should deny access to other user's plan."""
|
||||||
|
response = self.client1.get(
|
||||||
|
f"/api/installment-plans/{self.user2_installment.id}/"
|
||||||
|
)
|
||||||
|
|
||||||
|
self.assertIn(response.status_code, ACCESS_DENIED_CODES)
|
||||||
|
|
||||||
|
def test_user_cannot_see_other_users_recurring_transactions(self):
|
||||||
|
"""GET /api/recurring-transactions/ should only return user's own recurring."""
|
||||||
|
response = self.client1.get("/api/recurring-transactions/")
|
||||||
|
|
||||||
|
self.assertEqual(response.status_code, status.HTTP_200_OK)
|
||||||
|
|
||||||
|
recurring_ids = [r["id"] for r in response.data["results"]]
|
||||||
|
self.assertIn(self.user1_recurring.id, recurring_ids)
|
||||||
|
self.assertNotIn(self.user2_recurring.id, recurring_ids)
|
||||||
|
|
||||||
|
def test_user_cannot_access_other_users_recurring_transaction_detail(self):
|
||||||
|
"""GET /api/recurring-transactions/{id}/ should deny access to other user's recurring."""
|
||||||
|
response = self.client1.get(
|
||||||
|
f"/api/recurring-transactions/{self.user2_recurring.id}/"
|
||||||
|
)
|
||||||
|
|
||||||
|
self.assertIn(response.status_code, ACCESS_DENIED_CODES)
|
||||||
|
|
||||||
|
def test_user_cannot_modify_other_users_installment_plan(self):
|
||||||
|
"""PATCH on other user's installment plan should deny access."""
|
||||||
|
response = self.client1.patch(
|
||||||
|
f"/api/installment-plans/{self.user2_installment.id}/",
|
||||||
|
{"description": "Hacked Installment"},
|
||||||
|
)
|
||||||
|
|
||||||
|
self.assertIn(response.status_code, ACCESS_DENIED_CODES)
|
||||||
|
|
||||||
|
def test_user_cannot_delete_other_users_recurring_transaction(self):
|
||||||
|
"""DELETE on other user's recurring transaction should deny access."""
|
||||||
|
response = self.client1.delete(
|
||||||
|
f"/api/recurring-transactions/{self.user2_recurring.id}/"
|
||||||
|
)
|
||||||
|
|
||||||
|
self.assertIn(response.status_code, ACCESS_DENIED_CODES)
|
||||||
|
|
||||||
|
self.assertTrue(
|
||||||
|
RecurringTransaction.all_objects.filter(id=self.user2_recurring.id).exists()
|
||||||
|
)
|
||||||
@@ -0,0 +1,166 @@
|
|||||||
|
from datetime import timedelta
|
||||||
|
|
||||||
|
from django.contrib.auth import get_user_model
|
||||||
|
from django.test import TestCase, override_settings
|
||||||
|
from django.urls import reverse
|
||||||
|
from django.utils import timezone
|
||||||
|
from oauth2_provider.models import get_access_token_model, get_application_model
|
||||||
|
|
||||||
|
from apps.users.models import APIToken
|
||||||
|
|
||||||
|
User = get_user_model()
|
||||||
|
Application = get_application_model()
|
||||||
|
AccessToken = get_access_token_model()
|
||||||
|
|
||||||
|
|
||||||
|
@override_settings(DEMO=True)
|
||||||
|
class DemoModeAPITests(TestCase):
|
||||||
|
"""The DEMO-mode gate (apps.api.permissions.NotInDemoMode) must reject
|
||||||
|
API access regardless of the authentication method used, including the
|
||||||
|
PAT and OAuth2 backends introduced for MCP integrations."""
|
||||||
|
|
||||||
|
def setUp(self):
|
||||||
|
self.user = User.objects.create_user(
|
||||||
|
email="demo@example.com",
|
||||||
|
password="test-password",
|
||||||
|
)
|
||||||
|
|
||||||
|
def test_pat_cannot_access_api_in_demo_mode(self):
|
||||||
|
_token, raw_token = APIToken.objects.create_token(
|
||||||
|
user=self.user, name="n8n"
|
||||||
|
)
|
||||||
|
|
||||||
|
response = self.client.get(
|
||||||
|
"/api/accounts/",
|
||||||
|
HTTP_AUTHORIZATION=f"Token {raw_token}",
|
||||||
|
)
|
||||||
|
|
||||||
|
self.assertEqual(response.status_code, 403)
|
||||||
|
|
||||||
|
def test_oauth_access_token_cannot_access_api_in_demo_mode(self):
|
||||||
|
app = Application.objects.create(
|
||||||
|
name="Test Client",
|
||||||
|
client_type=Application.CLIENT_CONFIDENTIAL,
|
||||||
|
authorization_grant_type=Application.GRANT_AUTHORIZATION_CODE,
|
||||||
|
redirect_uris="http://127.0.0.1:8765/callback",
|
||||||
|
client_secret="secret",
|
||||||
|
)
|
||||||
|
access_token = AccessToken.objects.create(
|
||||||
|
user=self.user,
|
||||||
|
scope="mcp",
|
||||||
|
expires=timezone.now() + timedelta(hours=1),
|
||||||
|
token="demo-oauth-access-token-xyz",
|
||||||
|
application=app,
|
||||||
|
)
|
||||||
|
|
||||||
|
response = self.client.get(
|
||||||
|
"/api/accounts/",
|
||||||
|
HTTP_AUTHORIZATION=f"Bearer {access_token.token}",
|
||||||
|
)
|
||||||
|
|
||||||
|
self.assertEqual(response.status_code, 403)
|
||||||
|
|
||||||
|
def test_superuser_pat_can_access_api_in_demo_mode(self):
|
||||||
|
admin = User.objects.create_superuser(
|
||||||
|
email="admin@example.com",
|
||||||
|
password="test-password",
|
||||||
|
)
|
||||||
|
_token, raw_token = APIToken.objects.create_token(
|
||||||
|
user=admin, name="admin"
|
||||||
|
)
|
||||||
|
|
||||||
|
response = self.client.get(
|
||||||
|
"/api/accounts/",
|
||||||
|
HTTP_AUTHORIZATION=f"Token {raw_token}",
|
||||||
|
)
|
||||||
|
|
||||||
|
# NotInDemoMode grants superusers access in DEMO mode; the request is
|
||||||
|
# authenticated by the PAT, so the API responds normally (never 403).
|
||||||
|
self.assertNotEqual(response.status_code, 403)
|
||||||
|
|
||||||
|
|
||||||
|
@override_settings(DEMO=True)
|
||||||
|
class DemoModeOAuthEndpointTests(TestCase):
|
||||||
|
"""OAuth2 issuance and discovery endpoints must be disabled in DEMO mode
|
||||||
|
so demo tenants cannot obtain (or even discover) credentials."""
|
||||||
|
|
||||||
|
def setUp(self):
|
||||||
|
self.user = User.objects.create_user(
|
||||||
|
email="demo@example.com",
|
||||||
|
password="test-password",
|
||||||
|
)
|
||||||
|
|
||||||
|
def test_oauth_authorize_rejects_non_superuser_in_demo_mode(self):
|
||||||
|
self.client.force_login(self.user)
|
||||||
|
|
||||||
|
response = self.client.get(reverse("oauth2_provider:authorize"))
|
||||||
|
|
||||||
|
self.assertEqual(response.status_code, 403)
|
||||||
|
|
||||||
|
def test_oauth_token_rejects_non_superuser_in_demo_mode(self):
|
||||||
|
self.client.force_login(self.user)
|
||||||
|
|
||||||
|
response = self.client.post(reverse("oauth2_provider:token"))
|
||||||
|
|
||||||
|
self.assertEqual(response.status_code, 403)
|
||||||
|
|
||||||
|
def test_oauth_authorization_server_metadata_rejects_in_demo_mode(self):
|
||||||
|
response = self.client.get(reverse("oauth-authorization-server-metadata"))
|
||||||
|
|
||||||
|
self.assertEqual(response.status_code, 403)
|
||||||
|
|
||||||
|
def test_oauth_dynamic_client_registration_rejects_in_demo_mode(self):
|
||||||
|
response = self.client.post(
|
||||||
|
reverse("oauth-dynamic-client-registration"),
|
||||||
|
data="{}",
|
||||||
|
content_type="application/json",
|
||||||
|
)
|
||||||
|
|
||||||
|
self.assertEqual(response.status_code, 403)
|
||||||
|
|
||||||
|
|
||||||
|
@override_settings(DEMO=True)
|
||||||
|
class DemoModeAPITokenViewsTests(TestCase):
|
||||||
|
"""The PAT management UI must be disabled in DEMO mode just like the
|
||||||
|
other mutating user views."""
|
||||||
|
|
||||||
|
def setUp(self):
|
||||||
|
self.user = User.objects.create_user(
|
||||||
|
email="demo@example.com",
|
||||||
|
password="test-password",
|
||||||
|
)
|
||||||
|
self.client.force_login(self.user)
|
||||||
|
self.htmx_headers = {"HTTP_HX_REQUEST": "true"}
|
||||||
|
|
||||||
|
def test_cannot_create_api_token_from_ui_in_demo_mode(self):
|
||||||
|
response = self.client.post(
|
||||||
|
reverse("user_api_token_add"),
|
||||||
|
{"name": "n8n", "expires_in_days": "30"},
|
||||||
|
**self.htmx_headers,
|
||||||
|
)
|
||||||
|
|
||||||
|
self.assertEqual(response.status_code, 403)
|
||||||
|
self.assertEqual(APIToken.objects.count(), 0)
|
||||||
|
|
||||||
|
def test_cannot_revoke_api_token_from_ui_in_demo_mode(self):
|
||||||
|
token, _ = APIToken.objects.create_token(user=self.user, name="n8n")
|
||||||
|
|
||||||
|
response = self.client.delete(
|
||||||
|
reverse("user_api_token_revoke", kwargs={"token_id": token.id}),
|
||||||
|
**self.htmx_headers,
|
||||||
|
)
|
||||||
|
|
||||||
|
self.assertEqual(response.status_code, 403)
|
||||||
|
token.refresh_from_db()
|
||||||
|
self.assertIsNone(token.revoked_at)
|
||||||
|
|
||||||
|
def test_cannot_delete_api_token_from_ui_in_demo_mode(self):
|
||||||
|
token, _ = APIToken.objects.create_token(user=self.user, name="n8n")
|
||||||
|
|
||||||
|
response = self.client.delete(
|
||||||
|
reverse("user_api_token_delete", kwargs={"token_id": token.id}),
|
||||||
|
**self.htmx_headers,
|
||||||
|
)
|
||||||
|
|
||||||
|
self.assertEqual(response.status_code, 403)
|
||||||
|
self.assertTrue(APIToken.objects.filter(id=token.id).exists())
|
||||||
@@ -159,7 +159,7 @@ column_mapping:
|
|||||||
self.assertIn("import_run_id", response.data)
|
self.assertIn("import_run_id", response.data)
|
||||||
|
|
||||||
def test_unauthenticated_request(self):
|
def test_unauthenticated_request(self):
|
||||||
"""Test unauthenticated request returns 403"""
|
"""Test unauthenticated request returns 401"""
|
||||||
unauthenticated_client = APIClient()
|
unauthenticated_client = APIClient()
|
||||||
|
|
||||||
csv_content = b"date,description,amount\n2025-01-01,Test,100"
|
csv_content = b"date,description,amount\n2025-01-01,Test,100"
|
||||||
@@ -173,7 +173,7 @@ column_mapping:
|
|||||||
format="multipart",
|
format="multipart",
|
||||||
)
|
)
|
||||||
|
|
||||||
self.assertEqual(response.status_code, status.HTTP_403_FORBIDDEN)
|
self.assertEqual(response.status_code, status.HTTP_401_UNAUTHORIZED)
|
||||||
|
|
||||||
|
|
||||||
@override_settings(
|
@override_settings(
|
||||||
@@ -266,11 +266,11 @@ column_mapping:
|
|||||||
self.assertEqual(response.status_code, status.HTTP_404_NOT_FOUND)
|
self.assertEqual(response.status_code, status.HTTP_404_NOT_FOUND)
|
||||||
|
|
||||||
def test_profiles_unauthenticated(self):
|
def test_profiles_unauthenticated(self):
|
||||||
"""Test unauthenticated request returns 403"""
|
"""Test unauthenticated request returns 401"""
|
||||||
unauthenticated_client = APIClient()
|
unauthenticated_client = APIClient()
|
||||||
response = unauthenticated_client.get("/api/import/profiles/")
|
response = unauthenticated_client.get("/api/import/profiles/")
|
||||||
|
|
||||||
self.assertEqual(response.status_code, status.HTTP_403_FORBIDDEN)
|
self.assertEqual(response.status_code, status.HTTP_401_UNAUTHORIZED)
|
||||||
|
|
||||||
|
|
||||||
@override_settings(
|
@override_settings(
|
||||||
@@ -397,8 +397,8 @@ column_mapping:
|
|||||||
self.assertEqual(response.status_code, status.HTTP_404_NOT_FOUND)
|
self.assertEqual(response.status_code, status.HTTP_404_NOT_FOUND)
|
||||||
|
|
||||||
def test_runs_unauthenticated(self):
|
def test_runs_unauthenticated(self):
|
||||||
"""Test unauthenticated request returns 403"""
|
"""Test unauthenticated request returns 401"""
|
||||||
unauthenticated_client = APIClient()
|
unauthenticated_client = APIClient()
|
||||||
response = unauthenticated_client.get("/api/import/runs/")
|
response = unauthenticated_client.get("/api/import/runs/")
|
||||||
|
|
||||||
self.assertEqual(response.status_code, status.HTTP_403_FORBIDDEN)
|
self.assertEqual(response.status_code, status.HTTP_401_UNAUTHORIZED)
|
||||||
|
|||||||
@@ -0,0 +1,171 @@
|
|||||||
|
"""Object-level ownership on the SharedObject API viewsets.
|
||||||
|
|
||||||
|
DjangoModelPermissions is model-level: a user holding ``change_account`` could
|
||||||
|
write any object the viewset's queryset returned, and for SharedObject that
|
||||||
|
queryset includes other people's public and shared-with-them objects. Ordinary
|
||||||
|
users hold no model permissions, so this was not reachable for them, but the
|
||||||
|
object-level check was missing entirely.
|
||||||
|
|
||||||
|
Reads are unaffected -- they stay governed by SharedObjectManager.
|
||||||
|
"""
|
||||||
|
|
||||||
|
from datetime import date
|
||||||
|
from decimal import Decimal
|
||||||
|
|
||||||
|
from django.contrib.auth import get_user_model
|
||||||
|
from django.contrib.auth.models import Permission
|
||||||
|
from django.test import TestCase, override_settings
|
||||||
|
from rest_framework import status
|
||||||
|
from rest_framework.test import APIClient
|
||||||
|
|
||||||
|
from apps.accounts.models import Account
|
||||||
|
from apps.currencies.models import Currency
|
||||||
|
from apps.dca.models import DCAEntry, DCAStrategy
|
||||||
|
from apps.transactions.models import TransactionCategory
|
||||||
|
|
||||||
|
|
||||||
|
@override_settings(
|
||||||
|
STORAGES={
|
||||||
|
"default": {"BACKEND": "django.core.files.storage.FileSystemStorage"},
|
||||||
|
"staticfiles": {
|
||||||
|
"BACKEND": "django.contrib.staticfiles.storage.StaticFilesStorage"
|
||||||
|
},
|
||||||
|
},
|
||||||
|
WHITENOISE_AUTOREFRESH=True,
|
||||||
|
DEMO=False,
|
||||||
|
)
|
||||||
|
class SharedObjectAPIPermissionTests(TestCase):
|
||||||
|
"""The attacker here deliberately HOLDS the Django model permissions.
|
||||||
|
|
||||||
|
Without them DjangoModelPermissions already answers 403 and the
|
||||||
|
object-level check is never consulted, so the test would pass whether or
|
||||||
|
not it exists.
|
||||||
|
"""
|
||||||
|
|
||||||
|
def setUp(self):
|
||||||
|
User = get_user_model()
|
||||||
|
self.owner = User.objects.create_user(
|
||||||
|
email="owner@test.com", password="testpass123"
|
||||||
|
)
|
||||||
|
self.attacker = User.objects.create_user(
|
||||||
|
email="attacker@test.com", password="testpass123"
|
||||||
|
)
|
||||||
|
self.attacker.user_permissions.set(
|
||||||
|
Permission.objects.filter(
|
||||||
|
codename__in=[
|
||||||
|
"add_account",
|
||||||
|
"change_account",
|
||||||
|
"delete_account",
|
||||||
|
"add_transactioncategory",
|
||||||
|
"change_transactioncategory",
|
||||||
|
"delete_transactioncategory",
|
||||||
|
"add_dcaentry",
|
||||||
|
"change_dcaentry",
|
||||||
|
"delete_dcaentry",
|
||||||
|
]
|
||||||
|
)
|
||||||
|
)
|
||||||
|
# Permissions are cached on the user instance.
|
||||||
|
self.attacker = User.objects.get(pk=self.attacker.pk)
|
||||||
|
|
||||||
|
self.currency = Currency.objects.create(
|
||||||
|
code="USD", name="US Dollar", decimal_places=2
|
||||||
|
)
|
||||||
|
|
||||||
|
self.api = APIClient()
|
||||||
|
self.api.force_authenticate(user=self.attacker)
|
||||||
|
|
||||||
|
def test_cannot_modify_public_account(self):
|
||||||
|
account = Account.all_objects.create(
|
||||||
|
name="Public account",
|
||||||
|
currency=self.currency,
|
||||||
|
owner=self.owner,
|
||||||
|
visibility="public",
|
||||||
|
)
|
||||||
|
|
||||||
|
response = self.api.patch(
|
||||||
|
f"/api/accounts/{account.id}/", {"name": "HIJACKED"}, format="json"
|
||||||
|
)
|
||||||
|
|
||||||
|
self.assertEqual(response.status_code, status.HTTP_403_FORBIDDEN)
|
||||||
|
account.refresh_from_db()
|
||||||
|
self.assertEqual(account.name, "Public account")
|
||||||
|
|
||||||
|
def test_cannot_delete_account_shared_with_them(self):
|
||||||
|
account = Account.all_objects.create(
|
||||||
|
name="Shared account",
|
||||||
|
currency=self.currency,
|
||||||
|
owner=self.owner,
|
||||||
|
visibility="private",
|
||||||
|
)
|
||||||
|
account.shared_with.add(self.attacker)
|
||||||
|
|
||||||
|
response = self.api.delete(f"/api/accounts/{account.id}/")
|
||||||
|
|
||||||
|
self.assertEqual(response.status_code, status.HTTP_403_FORBIDDEN)
|
||||||
|
self.assertTrue(Account.all_objects.filter(pk=account.pk).exists())
|
||||||
|
|
||||||
|
def test_cannot_delete_public_category(self):
|
||||||
|
category = TransactionCategory.all_objects.create(
|
||||||
|
name="Public category", owner=self.owner, visibility="public"
|
||||||
|
)
|
||||||
|
|
||||||
|
response = self.api.delete(f"/api/categories/{category.id}/")
|
||||||
|
|
||||||
|
self.assertEqual(response.status_code, status.HTTP_403_FORBIDDEN)
|
||||||
|
self.assertTrue(TransactionCategory.all_objects.filter(pk=category.pk).exists())
|
||||||
|
|
||||||
|
def test_cannot_delete_entry_on_public_strategy(self):
|
||||||
|
strategy = DCAStrategy.all_objects.create(
|
||||||
|
name="Public strategy",
|
||||||
|
owner=self.owner,
|
||||||
|
visibility="public",
|
||||||
|
target_currency=self.currency,
|
||||||
|
payment_currency=self.currency,
|
||||||
|
)
|
||||||
|
entry = DCAEntry.objects.create(
|
||||||
|
strategy=strategy,
|
||||||
|
date=date(2025, 1, 1),
|
||||||
|
amount_paid=Decimal("100"),
|
||||||
|
amount_received=Decimal("1"),
|
||||||
|
)
|
||||||
|
|
||||||
|
response = self.api.delete(f"/api/dca/entries/{entry.id}/")
|
||||||
|
|
||||||
|
self.assertEqual(response.status_code, status.HTTP_403_FORBIDDEN)
|
||||||
|
self.assertTrue(DCAEntry.objects.filter(pk=entry.pk).exists())
|
||||||
|
|
||||||
|
def test_reads_of_shared_objects_still_work(self):
|
||||||
|
account = Account.all_objects.create(
|
||||||
|
name="Public account",
|
||||||
|
currency=self.currency,
|
||||||
|
owner=self.owner,
|
||||||
|
visibility="public",
|
||||||
|
)
|
||||||
|
|
||||||
|
response = self.api.get(f"/api/accounts/{account.id}/")
|
||||||
|
|
||||||
|
self.assertEqual(response.status_code, status.HTTP_200_OK)
|
||||||
|
|
||||||
|
def test_owner_can_still_write_their_own_objects(self):
|
||||||
|
owner_client = APIClient()
|
||||||
|
self.owner.user_permissions.set(
|
||||||
|
Permission.objects.filter(codename__in=["change_account", "delete_account"])
|
||||||
|
)
|
||||||
|
owner = get_user_model().objects.get(pk=self.owner.pk)
|
||||||
|
owner_client.force_authenticate(user=owner)
|
||||||
|
|
||||||
|
account = Account.all_objects.create(
|
||||||
|
name="Own account",
|
||||||
|
currency=self.currency,
|
||||||
|
owner=owner,
|
||||||
|
visibility="private",
|
||||||
|
)
|
||||||
|
|
||||||
|
response = owner_client.patch(
|
||||||
|
f"/api/accounts/{account.id}/", {"name": "Renamed"}, format="json"
|
||||||
|
)
|
||||||
|
|
||||||
|
self.assertEqual(response.status_code, status.HTTP_200_OK)
|
||||||
|
account.refresh_from_db()
|
||||||
|
self.assertEqual(account.name, "Renamed")
|
||||||
@@ -0,0 +1,587 @@
|
|||||||
|
from datetime import date
|
||||||
|
from decimal import Decimal
|
||||||
|
|
||||||
|
from django.contrib.auth import get_user_model
|
||||||
|
from django.test import TestCase, override_settings
|
||||||
|
from rest_framework import status
|
||||||
|
from rest_framework.test import APIClient
|
||||||
|
|
||||||
|
from apps.accounts.models import Account, AccountGroup
|
||||||
|
from apps.currencies.models import Currency
|
||||||
|
from apps.dca.models import DCAStrategy, DCAEntry
|
||||||
|
from apps.transactions.models import (
|
||||||
|
Transaction,
|
||||||
|
TransactionCategory,
|
||||||
|
TransactionTag,
|
||||||
|
TransactionEntity,
|
||||||
|
)
|
||||||
|
|
||||||
|
|
||||||
|
ACCESS_DENIED_CODES = [status.HTTP_403_FORBIDDEN, status.HTTP_404_NOT_FOUND]
|
||||||
|
|
||||||
|
|
||||||
|
@override_settings(
|
||||||
|
STORAGES={
|
||||||
|
"default": {"BACKEND": "django.core.files.storage.FileSystemStorage"},
|
||||||
|
"staticfiles": {
|
||||||
|
"BACKEND": "django.contrib.staticfiles.storage.StaticFilesStorage"
|
||||||
|
},
|
||||||
|
},
|
||||||
|
WHITENOISE_AUTOREFRESH=True,
|
||||||
|
)
|
||||||
|
class SharedAccountAccessTests(TestCase):
|
||||||
|
"""Tests for shared account access via shared_with field."""
|
||||||
|
|
||||||
|
def setUp(self):
|
||||||
|
"""Set up test data with shared accounts."""
|
||||||
|
User = get_user_model()
|
||||||
|
|
||||||
|
# User 1 - owner
|
||||||
|
self.user1 = User.objects.create_user(
|
||||||
|
email="user1@test.com", password="testpass123"
|
||||||
|
)
|
||||||
|
self.client1 = APIClient()
|
||||||
|
self.client1.force_authenticate(user=self.user1)
|
||||||
|
|
||||||
|
# User 2 - will have shared access
|
||||||
|
self.user2 = User.objects.create_user(
|
||||||
|
email="user2@test.com", password="testpass123"
|
||||||
|
)
|
||||||
|
self.client2 = APIClient()
|
||||||
|
self.client2.force_authenticate(user=self.user2)
|
||||||
|
|
||||||
|
# User 3 - no shared access
|
||||||
|
self.user3 = User.objects.create_user(
|
||||||
|
email="user3@test.com", password="testpass123"
|
||||||
|
)
|
||||||
|
self.client3 = APIClient()
|
||||||
|
self.client3.force_authenticate(user=self.user3)
|
||||||
|
|
||||||
|
self.currency = Currency.objects.create(
|
||||||
|
code="USD", name="US Dollar", decimal_places=2, prefix="$ "
|
||||||
|
)
|
||||||
|
|
||||||
|
# User 1's account shared with user 2
|
||||||
|
self.shared_account = Account.all_objects.create(
|
||||||
|
name="Shared Account",
|
||||||
|
currency=self.currency,
|
||||||
|
owner=self.user1,
|
||||||
|
visibility="private",
|
||||||
|
)
|
||||||
|
self.shared_account.shared_with.add(self.user2)
|
||||||
|
|
||||||
|
# User 1's private account (not shared)
|
||||||
|
self.private_account = Account.all_objects.create(
|
||||||
|
name="Private Account",
|
||||||
|
currency=self.currency,
|
||||||
|
owner=self.user1,
|
||||||
|
visibility="private",
|
||||||
|
)
|
||||||
|
|
||||||
|
# Transaction in shared account
|
||||||
|
self.shared_transaction = Transaction.userless_all_objects.create(
|
||||||
|
account=self.shared_account,
|
||||||
|
type=Transaction.Type.INCOME,
|
||||||
|
amount=Decimal("100.00"),
|
||||||
|
is_paid=True,
|
||||||
|
date=date(2025, 1, 1),
|
||||||
|
description="Shared Transaction",
|
||||||
|
owner=self.user1,
|
||||||
|
)
|
||||||
|
|
||||||
|
# Transaction in private account
|
||||||
|
self.private_transaction = Transaction.userless_all_objects.create(
|
||||||
|
account=self.private_account,
|
||||||
|
type=Transaction.Type.EXPENSE,
|
||||||
|
amount=Decimal("50.00"),
|
||||||
|
is_paid=True,
|
||||||
|
date=date(2025, 1, 1),
|
||||||
|
description="Private Transaction",
|
||||||
|
owner=self.user1,
|
||||||
|
)
|
||||||
|
|
||||||
|
def test_user_can_see_accounts_shared_with_them(self):
|
||||||
|
"""User2 should see the account shared with them."""
|
||||||
|
response = self.client2.get("/api/accounts/")
|
||||||
|
|
||||||
|
self.assertEqual(response.status_code, status.HTTP_200_OK)
|
||||||
|
|
||||||
|
account_ids = [acc["id"] for acc in response.data["results"]]
|
||||||
|
self.assertIn(self.shared_account.id, account_ids)
|
||||||
|
|
||||||
|
def test_user_cannot_see_accounts_not_shared_with_them(self):
|
||||||
|
"""User2 should NOT see user1's private (non-shared) account."""
|
||||||
|
response = self.client2.get("/api/accounts/")
|
||||||
|
|
||||||
|
self.assertEqual(response.status_code, status.HTTP_200_OK)
|
||||||
|
|
||||||
|
account_ids = [acc["id"] for acc in response.data["results"]]
|
||||||
|
self.assertNotIn(self.private_account.id, account_ids)
|
||||||
|
|
||||||
|
def test_user_can_access_shared_account_detail(self):
|
||||||
|
"""User2 should be able to access shared account details."""
|
||||||
|
response = self.client2.get(f"/api/accounts/{self.shared_account.id}/")
|
||||||
|
|
||||||
|
self.assertEqual(response.status_code, status.HTTP_200_OK)
|
||||||
|
self.assertEqual(response.data["name"], "Shared Account")
|
||||||
|
|
||||||
|
def test_user_without_share_cannot_access_shared_account(self):
|
||||||
|
"""User3 should NOT be able to access the shared account."""
|
||||||
|
response = self.client3.get(f"/api/accounts/{self.shared_account.id}/")
|
||||||
|
|
||||||
|
self.assertIn(response.status_code, ACCESS_DENIED_CODES)
|
||||||
|
|
||||||
|
def test_user_can_see_transactions_in_shared_account(self):
|
||||||
|
"""User2 should see transactions in the shared account."""
|
||||||
|
response = self.client2.get("/api/transactions/")
|
||||||
|
|
||||||
|
self.assertEqual(response.status_code, status.HTTP_200_OK)
|
||||||
|
|
||||||
|
transaction_ids = [t["id"] for t in response.data["results"]]
|
||||||
|
self.assertIn(self.shared_transaction.id, transaction_ids)
|
||||||
|
self.assertNotIn(self.private_transaction.id, transaction_ids)
|
||||||
|
|
||||||
|
def test_user_can_access_transaction_in_shared_account(self):
|
||||||
|
"""User2 should be able to access transaction details in shared account."""
|
||||||
|
response = self.client2.get(f"/api/transactions/{self.shared_transaction.id}/")
|
||||||
|
|
||||||
|
self.assertEqual(response.status_code, status.HTTP_200_OK)
|
||||||
|
self.assertEqual(response.data["description"], "Shared Transaction")
|
||||||
|
|
||||||
|
def test_user_cannot_access_transaction_in_non_shared_account(self):
|
||||||
|
"""User2 should NOT access transactions in user1's private account."""
|
||||||
|
response = self.client2.get(f"/api/transactions/{self.private_transaction.id}/")
|
||||||
|
|
||||||
|
self.assertIn(response.status_code, ACCESS_DENIED_CODES)
|
||||||
|
|
||||||
|
def test_user_can_get_balance_of_shared_account(self):
|
||||||
|
"""User2 should be able to get balance of shared account."""
|
||||||
|
response = self.client2.get(f"/api/accounts/{self.shared_account.id}/balance/")
|
||||||
|
|
||||||
|
self.assertEqual(response.status_code, status.HTTP_200_OK)
|
||||||
|
self.assertIn("current_balance", response.data)
|
||||||
|
|
||||||
|
def test_sharing_works_with_multiple_users(self):
|
||||||
|
"""Account shared with multiple users should be accessible by all."""
|
||||||
|
# Add user3 to shared_with
|
||||||
|
self.shared_account.shared_with.add(self.user3)
|
||||||
|
|
||||||
|
# User2 still has access
|
||||||
|
response2 = self.client2.get(f"/api/accounts/{self.shared_account.id}/")
|
||||||
|
self.assertEqual(response2.status_code, status.HTTP_200_OK)
|
||||||
|
|
||||||
|
# User3 now has access
|
||||||
|
response3 = self.client3.get(f"/api/accounts/{self.shared_account.id}/")
|
||||||
|
self.assertEqual(response3.status_code, status.HTTP_200_OK)
|
||||||
|
|
||||||
|
|
||||||
|
@override_settings(
|
||||||
|
STORAGES={
|
||||||
|
"default": {"BACKEND": "django.core.files.storage.FileSystemStorage"},
|
||||||
|
"staticfiles": {
|
||||||
|
"BACKEND": "django.contrib.staticfiles.storage.StaticFilesStorage"
|
||||||
|
},
|
||||||
|
},
|
||||||
|
WHITENOISE_AUTOREFRESH=True,
|
||||||
|
)
|
||||||
|
class PublicVisibilityTests(TestCase):
|
||||||
|
"""Tests for public visibility access."""
|
||||||
|
|
||||||
|
def setUp(self):
|
||||||
|
"""Set up test data with public accounts."""
|
||||||
|
User = get_user_model()
|
||||||
|
|
||||||
|
self.user1 = User.objects.create_user(
|
||||||
|
email="user1@test.com", password="testpass123"
|
||||||
|
)
|
||||||
|
self.client1 = APIClient()
|
||||||
|
self.client1.force_authenticate(user=self.user1)
|
||||||
|
|
||||||
|
self.user2 = User.objects.create_user(
|
||||||
|
email="user2@test.com", password="testpass123"
|
||||||
|
)
|
||||||
|
self.client2 = APIClient()
|
||||||
|
self.client2.force_authenticate(user=self.user2)
|
||||||
|
|
||||||
|
self.currency = Currency.objects.create(
|
||||||
|
code="USD", name="US Dollar", decimal_places=2, prefix="$ "
|
||||||
|
)
|
||||||
|
|
||||||
|
# User 1's public account
|
||||||
|
self.public_account = Account.all_objects.create(
|
||||||
|
name="Public Account",
|
||||||
|
currency=self.currency,
|
||||||
|
owner=self.user1,
|
||||||
|
visibility="public",
|
||||||
|
)
|
||||||
|
|
||||||
|
# User 1's private account
|
||||||
|
self.private_account = Account.all_objects.create(
|
||||||
|
name="Private Account",
|
||||||
|
currency=self.currency,
|
||||||
|
owner=self.user1,
|
||||||
|
visibility="private",
|
||||||
|
)
|
||||||
|
|
||||||
|
# Transaction in public account
|
||||||
|
self.public_transaction = Transaction.userless_all_objects.create(
|
||||||
|
account=self.public_account,
|
||||||
|
type=Transaction.Type.INCOME,
|
||||||
|
amount=Decimal("100.00"),
|
||||||
|
is_paid=True,
|
||||||
|
date=date(2025, 1, 1),
|
||||||
|
description="Public Transaction",
|
||||||
|
owner=self.user1,
|
||||||
|
)
|
||||||
|
|
||||||
|
def test_user_can_see_public_accounts(self):
|
||||||
|
"""User2 should see user1's public account."""
|
||||||
|
response = self.client2.get("/api/accounts/")
|
||||||
|
|
||||||
|
self.assertEqual(response.status_code, status.HTTP_200_OK)
|
||||||
|
|
||||||
|
account_ids = [acc["id"] for acc in response.data["results"]]
|
||||||
|
self.assertIn(self.public_account.id, account_ids)
|
||||||
|
self.assertNotIn(self.private_account.id, account_ids)
|
||||||
|
|
||||||
|
def test_user_can_access_public_account_detail(self):
|
||||||
|
"""User2 should be able to access public account details."""
|
||||||
|
response = self.client2.get(f"/api/accounts/{self.public_account.id}/")
|
||||||
|
|
||||||
|
self.assertEqual(response.status_code, status.HTTP_200_OK)
|
||||||
|
self.assertEqual(response.data["name"], "Public Account")
|
||||||
|
|
||||||
|
def test_user_can_see_transactions_in_public_accounts(self):
|
||||||
|
"""User2 should see transactions in public accounts."""
|
||||||
|
response = self.client2.get("/api/transactions/")
|
||||||
|
|
||||||
|
self.assertEqual(response.status_code, status.HTTP_200_OK)
|
||||||
|
|
||||||
|
transaction_ids = [t["id"] for t in response.data["results"]]
|
||||||
|
self.assertIn(self.public_transaction.id, transaction_ids)
|
||||||
|
|
||||||
|
|
||||||
|
@override_settings(
|
||||||
|
STORAGES={
|
||||||
|
"default": {"BACKEND": "django.core.files.storage.FileSystemStorage"},
|
||||||
|
"staticfiles": {
|
||||||
|
"BACKEND": "django.contrib.staticfiles.storage.StaticFilesStorage"
|
||||||
|
},
|
||||||
|
},
|
||||||
|
WHITENOISE_AUTOREFRESH=True,
|
||||||
|
)
|
||||||
|
class SharedCategoryTagEntityTests(TestCase):
|
||||||
|
"""Tests for shared categories, tags, and entities."""
|
||||||
|
|
||||||
|
def setUp(self):
|
||||||
|
"""Set up test data with shared categories/tags/entities."""
|
||||||
|
User = get_user_model()
|
||||||
|
|
||||||
|
self.user1 = User.objects.create_user(
|
||||||
|
email="user1@test.com", password="testpass123"
|
||||||
|
)
|
||||||
|
self.client1 = APIClient()
|
||||||
|
self.client1.force_authenticate(user=self.user1)
|
||||||
|
|
||||||
|
self.user2 = User.objects.create_user(
|
||||||
|
email="user2@test.com", password="testpass123"
|
||||||
|
)
|
||||||
|
self.client2 = APIClient()
|
||||||
|
self.client2.force_authenticate(user=self.user2)
|
||||||
|
|
||||||
|
self.user3 = User.objects.create_user(
|
||||||
|
email="user3@test.com", password="testpass123"
|
||||||
|
)
|
||||||
|
self.client3 = APIClient()
|
||||||
|
self.client3.force_authenticate(user=self.user3)
|
||||||
|
|
||||||
|
# User 1's category shared with user 2
|
||||||
|
self.shared_category = TransactionCategory.all_objects.create(
|
||||||
|
name="Shared Category", owner=self.user1
|
||||||
|
)
|
||||||
|
self.shared_category.shared_with.add(self.user2)
|
||||||
|
|
||||||
|
# User 1's private category
|
||||||
|
self.private_category = TransactionCategory.all_objects.create(
|
||||||
|
name="Private Category", owner=self.user1
|
||||||
|
)
|
||||||
|
|
||||||
|
# User 1's public category
|
||||||
|
self.public_category = TransactionCategory.all_objects.create(
|
||||||
|
name="Public Category", owner=self.user1, visibility="public"
|
||||||
|
)
|
||||||
|
|
||||||
|
# User 1's tag shared with user 2
|
||||||
|
self.shared_tag = TransactionTag.all_objects.create(
|
||||||
|
name="Shared Tag", owner=self.user1
|
||||||
|
)
|
||||||
|
self.shared_tag.shared_with.add(self.user2)
|
||||||
|
|
||||||
|
# User 1's entity shared with user 2
|
||||||
|
self.shared_entity = TransactionEntity.all_objects.create(
|
||||||
|
name="Shared Entity", owner=self.user1
|
||||||
|
)
|
||||||
|
self.shared_entity.shared_with.add(self.user2)
|
||||||
|
|
||||||
|
def test_user_can_see_shared_categories(self):
|
||||||
|
"""User2 should see categories shared with them."""
|
||||||
|
response = self.client2.get("/api/categories/")
|
||||||
|
|
||||||
|
self.assertEqual(response.status_code, status.HTTP_200_OK)
|
||||||
|
|
||||||
|
category_ids = [c["id"] for c in response.data["results"]]
|
||||||
|
self.assertIn(self.shared_category.id, category_ids)
|
||||||
|
self.assertNotIn(self.private_category.id, category_ids)
|
||||||
|
|
||||||
|
def test_user_can_access_shared_category_detail(self):
|
||||||
|
"""User2 should be able to access shared category details."""
|
||||||
|
response = self.client2.get(f"/api/categories/{self.shared_category.id}/")
|
||||||
|
|
||||||
|
self.assertEqual(response.status_code, status.HTTP_200_OK)
|
||||||
|
self.assertEqual(response.data["name"], "Shared Category")
|
||||||
|
|
||||||
|
def test_user_can_see_public_categories(self):
|
||||||
|
"""User3 should see public categories."""
|
||||||
|
response = self.client3.get("/api/categories/")
|
||||||
|
|
||||||
|
self.assertEqual(response.status_code, status.HTTP_200_OK)
|
||||||
|
|
||||||
|
category_ids = [c["id"] for c in response.data["results"]]
|
||||||
|
self.assertIn(self.public_category.id, category_ids)
|
||||||
|
|
||||||
|
def test_user_without_share_cannot_see_shared_category(self):
|
||||||
|
"""User3 should NOT see category shared only with user2."""
|
||||||
|
response = self.client3.get("/api/categories/")
|
||||||
|
|
||||||
|
self.assertEqual(response.status_code, status.HTTP_200_OK)
|
||||||
|
|
||||||
|
category_ids = [c["id"] for c in response.data["results"]]
|
||||||
|
self.assertNotIn(self.shared_category.id, category_ids)
|
||||||
|
|
||||||
|
def test_user_can_see_shared_tags(self):
|
||||||
|
"""User2 should see tags shared with them."""
|
||||||
|
response = self.client2.get("/api/tags/")
|
||||||
|
|
||||||
|
self.assertEqual(response.status_code, status.HTTP_200_OK)
|
||||||
|
|
||||||
|
tag_ids = [t["id"] for t in response.data["results"]]
|
||||||
|
self.assertIn(self.shared_tag.id, tag_ids)
|
||||||
|
|
||||||
|
def test_user_can_access_shared_tag_detail(self):
|
||||||
|
"""User2 should be able to access shared tag details."""
|
||||||
|
response = self.client2.get(f"/api/tags/{self.shared_tag.id}/")
|
||||||
|
|
||||||
|
self.assertEqual(response.status_code, status.HTTP_200_OK)
|
||||||
|
self.assertEqual(response.data["name"], "Shared Tag")
|
||||||
|
|
||||||
|
def test_user_can_see_shared_entities(self):
|
||||||
|
"""User2 should see entities shared with them."""
|
||||||
|
response = self.client2.get("/api/entities/")
|
||||||
|
|
||||||
|
self.assertEqual(response.status_code, status.HTTP_200_OK)
|
||||||
|
|
||||||
|
entity_ids = [e["id"] for e in response.data["results"]]
|
||||||
|
self.assertIn(self.shared_entity.id, entity_ids)
|
||||||
|
|
||||||
|
def test_user_can_access_shared_entity_detail(self):
|
||||||
|
"""User2 should be able to access shared entity details."""
|
||||||
|
response = self.client2.get(f"/api/entities/{self.shared_entity.id}/")
|
||||||
|
|
||||||
|
self.assertEqual(response.status_code, status.HTTP_200_OK)
|
||||||
|
self.assertEqual(response.data["name"], "Shared Entity")
|
||||||
|
|
||||||
|
|
||||||
|
@override_settings(
|
||||||
|
STORAGES={
|
||||||
|
"default": {"BACKEND": "django.core.files.storage.FileSystemStorage"},
|
||||||
|
"staticfiles": {
|
||||||
|
"BACKEND": "django.contrib.staticfiles.storage.StaticFilesStorage"
|
||||||
|
},
|
||||||
|
},
|
||||||
|
WHITENOISE_AUTOREFRESH=True,
|
||||||
|
)
|
||||||
|
class SharedDCAAccessTests(TestCase):
|
||||||
|
"""Tests for shared DCA strategy access."""
|
||||||
|
|
||||||
|
def setUp(self):
|
||||||
|
"""Set up test data with shared DCA strategies."""
|
||||||
|
User = get_user_model()
|
||||||
|
|
||||||
|
self.user1 = User.objects.create_user(
|
||||||
|
email="user1@test.com", password="testpass123"
|
||||||
|
)
|
||||||
|
self.client1 = APIClient()
|
||||||
|
self.client1.force_authenticate(user=self.user1)
|
||||||
|
|
||||||
|
self.user2 = User.objects.create_user(
|
||||||
|
email="user2@test.com", password="testpass123"
|
||||||
|
)
|
||||||
|
self.client2 = APIClient()
|
||||||
|
self.client2.force_authenticate(user=self.user2)
|
||||||
|
|
||||||
|
self.user3 = User.objects.create_user(
|
||||||
|
email="user3@test.com", password="testpass123"
|
||||||
|
)
|
||||||
|
self.client3 = APIClient()
|
||||||
|
self.client3.force_authenticate(user=self.user3)
|
||||||
|
|
||||||
|
self.currency1 = Currency.objects.create(
|
||||||
|
code="BTC", name="Bitcoin", decimal_places=8, prefix=""
|
||||||
|
)
|
||||||
|
self.currency2 = Currency.objects.create(
|
||||||
|
code="USD", name="US Dollar", decimal_places=2, prefix="$ "
|
||||||
|
)
|
||||||
|
|
||||||
|
# User 1's DCA strategy shared with user 2
|
||||||
|
self.shared_strategy = DCAStrategy.all_objects.create(
|
||||||
|
name="Shared BTC Strategy",
|
||||||
|
target_currency=self.currency1,
|
||||||
|
payment_currency=self.currency2,
|
||||||
|
owner=self.user1,
|
||||||
|
)
|
||||||
|
self.shared_strategy.shared_with.add(self.user2)
|
||||||
|
|
||||||
|
# Entry in shared strategy
|
||||||
|
self.shared_entry = DCAEntry.objects.create(
|
||||||
|
strategy=self.shared_strategy,
|
||||||
|
date=date(2025, 1, 1),
|
||||||
|
amount_paid=Decimal("100.00"),
|
||||||
|
amount_received=Decimal("0.001"),
|
||||||
|
)
|
||||||
|
|
||||||
|
# User 1's private strategy
|
||||||
|
self.private_strategy = DCAStrategy.all_objects.create(
|
||||||
|
name="Private BTC Strategy",
|
||||||
|
target_currency=self.currency1,
|
||||||
|
payment_currency=self.currency2,
|
||||||
|
owner=self.user1,
|
||||||
|
)
|
||||||
|
|
||||||
|
def test_user_can_see_shared_dca_strategies(self):
|
||||||
|
"""User2 should see DCA strategies shared with them."""
|
||||||
|
response = self.client2.get("/api/dca/strategies/")
|
||||||
|
|
||||||
|
self.assertEqual(response.status_code, status.HTTP_200_OK)
|
||||||
|
|
||||||
|
strategy_ids = [s["id"] for s in response.data["results"]]
|
||||||
|
self.assertIn(self.shared_strategy.id, strategy_ids)
|
||||||
|
self.assertNotIn(self.private_strategy.id, strategy_ids)
|
||||||
|
|
||||||
|
def test_user_can_access_shared_dca_strategy_detail(self):
|
||||||
|
"""User2 should be able to access shared strategy details."""
|
||||||
|
response = self.client2.get(f"/api/dca/strategies/{self.shared_strategy.id}/")
|
||||||
|
|
||||||
|
self.assertEqual(response.status_code, status.HTTP_200_OK)
|
||||||
|
self.assertEqual(response.data["name"], "Shared BTC Strategy")
|
||||||
|
|
||||||
|
def test_user_without_share_cannot_see_shared_strategy(self):
|
||||||
|
"""User3 should NOT see strategy shared only with user2."""
|
||||||
|
response = self.client3.get("/api/dca/strategies/")
|
||||||
|
|
||||||
|
self.assertEqual(response.status_code, status.HTTP_200_OK)
|
||||||
|
|
||||||
|
strategy_ids = [s["id"] for s in response.data["results"]]
|
||||||
|
self.assertNotIn(self.shared_strategy.id, strategy_ids)
|
||||||
|
|
||||||
|
def test_user_can_access_shared_strategy_actions(self):
|
||||||
|
"""User2 should be able to access actions on shared strategy."""
|
||||||
|
# investment_frequency
|
||||||
|
response1 = self.client2.get(
|
||||||
|
f"/api/dca/strategies/{self.shared_strategy.id}/investment_frequency/"
|
||||||
|
)
|
||||||
|
self.assertEqual(response1.status_code, status.HTTP_200_OK)
|
||||||
|
|
||||||
|
# price_comparison
|
||||||
|
response2 = self.client2.get(
|
||||||
|
f"/api/dca/strategies/{self.shared_strategy.id}/price_comparison/"
|
||||||
|
)
|
||||||
|
self.assertEqual(response2.status_code, status.HTTP_200_OK)
|
||||||
|
|
||||||
|
# current_price
|
||||||
|
response3 = self.client2.get(
|
||||||
|
f"/api/dca/strategies/{self.shared_strategy.id}/current_price/"
|
||||||
|
)
|
||||||
|
self.assertEqual(response3.status_code, status.HTTP_200_OK)
|
||||||
|
|
||||||
|
|
||||||
|
@override_settings(
|
||||||
|
STORAGES={
|
||||||
|
"default": {"BACKEND": "django.core.files.storage.FileSystemStorage"},
|
||||||
|
"staticfiles": {
|
||||||
|
"BACKEND": "django.contrib.staticfiles.storage.StaticFilesStorage"
|
||||||
|
},
|
||||||
|
},
|
||||||
|
WHITENOISE_AUTOREFRESH=True,
|
||||||
|
)
|
||||||
|
class SharedAccountGroupTests(TestCase):
|
||||||
|
"""Tests for shared account group access."""
|
||||||
|
|
||||||
|
def setUp(self):
|
||||||
|
"""Set up test data with shared account groups."""
|
||||||
|
User = get_user_model()
|
||||||
|
|
||||||
|
self.user1 = User.objects.create_user(
|
||||||
|
email="user1@test.com", password="testpass123"
|
||||||
|
)
|
||||||
|
self.client1 = APIClient()
|
||||||
|
self.client1.force_authenticate(user=self.user1)
|
||||||
|
|
||||||
|
self.user2 = User.objects.create_user(
|
||||||
|
email="user2@test.com", password="testpass123"
|
||||||
|
)
|
||||||
|
self.client2 = APIClient()
|
||||||
|
self.client2.force_authenticate(user=self.user2)
|
||||||
|
|
||||||
|
self.user3 = User.objects.create_user(
|
||||||
|
email="user3@test.com", password="testpass123"
|
||||||
|
)
|
||||||
|
self.client3 = APIClient()
|
||||||
|
self.client3.force_authenticate(user=self.user3)
|
||||||
|
|
||||||
|
# User 1's account group shared with user 2
|
||||||
|
self.shared_group = AccountGroup.all_objects.create(
|
||||||
|
name="Shared Group", owner=self.user1
|
||||||
|
)
|
||||||
|
self.shared_group.shared_with.add(self.user2)
|
||||||
|
|
||||||
|
# User 1's private account group
|
||||||
|
self.private_group = AccountGroup.all_objects.create(
|
||||||
|
name="Private Group", owner=self.user1
|
||||||
|
)
|
||||||
|
|
||||||
|
# User 1's public account group
|
||||||
|
self.public_group = AccountGroup.all_objects.create(
|
||||||
|
name="Public Group", owner=self.user1, visibility="public"
|
||||||
|
)
|
||||||
|
|
||||||
|
def test_user_can_see_shared_account_groups(self):
|
||||||
|
"""User2 should see account groups shared with them."""
|
||||||
|
response = self.client2.get("/api/account-groups/")
|
||||||
|
|
||||||
|
self.assertEqual(response.status_code, status.HTTP_200_OK)
|
||||||
|
|
||||||
|
group_ids = [g["id"] for g in response.data["results"]]
|
||||||
|
self.assertIn(self.shared_group.id, group_ids)
|
||||||
|
self.assertNotIn(self.private_group.id, group_ids)
|
||||||
|
|
||||||
|
def test_user_can_access_shared_account_group_detail(self):
|
||||||
|
"""User2 should be able to access shared account group details."""
|
||||||
|
response = self.client2.get(f"/api/account-groups/{self.shared_group.id}/")
|
||||||
|
|
||||||
|
self.assertEqual(response.status_code, status.HTTP_200_OK)
|
||||||
|
self.assertEqual(response.data["name"], "Shared Group")
|
||||||
|
|
||||||
|
def test_user_can_see_public_account_groups(self):
|
||||||
|
"""User3 should see public account groups."""
|
||||||
|
response = self.client3.get("/api/account-groups/")
|
||||||
|
|
||||||
|
self.assertEqual(response.status_code, status.HTTP_200_OK)
|
||||||
|
|
||||||
|
group_ids = [g["id"] for g in response.data["results"]]
|
||||||
|
self.assertIn(self.public_group.id, group_ids)
|
||||||
|
|
||||||
|
def test_user_without_share_cannot_access_shared_group(self):
|
||||||
|
"""User3 should NOT be able to access shared account group."""
|
||||||
|
response = self.client3.get(f"/api/account-groups/{self.shared_group.id}/")
|
||||||
|
|
||||||
|
self.assertIn(response.status_code, ACCESS_DENIED_CODES)
|
||||||
@@ -6,19 +6,30 @@ from rest_framework.response import Response
|
|||||||
|
|
||||||
from apps.accounts.models import AccountGroup, Account
|
from apps.accounts.models import AccountGroup, Account
|
||||||
from apps.accounts.services import get_account_balance
|
from apps.accounts.services import get_account_balance
|
||||||
from apps.api.custom.pagination import CustomPageNumberPagination
|
from apps.api.permissions import SHARED_OBJECT_PERMISSIONS
|
||||||
from apps.api.serializers import AccountGroupSerializer, AccountSerializer, AccountBalanceSerializer
|
from apps.api.serializers import (
|
||||||
|
AccountGroupSerializer,
|
||||||
|
AccountSerializer,
|
||||||
|
AccountBalanceSerializer,
|
||||||
|
)
|
||||||
|
|
||||||
|
|
||||||
class AccountGroupViewSet(viewsets.ModelViewSet):
|
class AccountGroupViewSet(viewsets.ModelViewSet):
|
||||||
"""ViewSet for managing account groups."""
|
"""ViewSet for managing account groups."""
|
||||||
|
|
||||||
|
permission_classes = SHARED_OBJECT_PERMISSIONS
|
||||||
queryset = AccountGroup.objects.all()
|
queryset = AccountGroup.objects.all()
|
||||||
serializer_class = AccountGroupSerializer
|
serializer_class = AccountGroupSerializer
|
||||||
pagination_class = CustomPageNumberPagination
|
filterset_fields = {
|
||||||
|
"name": ["exact", "icontains"],
|
||||||
|
"owner": ["exact"],
|
||||||
|
}
|
||||||
|
search_fields = ["name"]
|
||||||
|
ordering_fields = "__all__"
|
||||||
|
ordering = ["id"]
|
||||||
|
|
||||||
def get_queryset(self):
|
def get_queryset(self):
|
||||||
return AccountGroup.objects.all().order_by("id")
|
return AccountGroup.objects.all()
|
||||||
|
|
||||||
|
|
||||||
@extend_schema_view(
|
@extend_schema_view(
|
||||||
@@ -31,30 +42,41 @@ class AccountGroupViewSet(viewsets.ModelViewSet):
|
|||||||
class AccountViewSet(viewsets.ModelViewSet):
|
class AccountViewSet(viewsets.ModelViewSet):
|
||||||
"""ViewSet for managing accounts."""
|
"""ViewSet for managing accounts."""
|
||||||
|
|
||||||
|
permission_classes = SHARED_OBJECT_PERMISSIONS
|
||||||
queryset = Account.objects.all()
|
queryset = Account.objects.all()
|
||||||
serializer_class = AccountSerializer
|
serializer_class = AccountSerializer
|
||||||
pagination_class = CustomPageNumberPagination
|
filterset_fields = {
|
||||||
|
"name": ["exact", "icontains"],
|
||||||
|
"group": ["exact", "isnull"],
|
||||||
|
"currency": ["exact"],
|
||||||
|
"exchange_currency": ["exact", "isnull"],
|
||||||
|
"is_asset": ["exact"],
|
||||||
|
"is_archived": ["exact"],
|
||||||
|
"owner": ["exact"],
|
||||||
|
}
|
||||||
|
search_fields = ["name"]
|
||||||
|
ordering_fields = "__all__"
|
||||||
|
ordering = ["id"]
|
||||||
|
|
||||||
def get_queryset(self):
|
def get_queryset(self):
|
||||||
return (
|
return Account.objects.all().select_related(
|
||||||
Account.objects.all()
|
"group", "currency", "exchange_currency"
|
||||||
.order_by("id")
|
|
||||||
.select_related("group", "currency", "exchange_currency")
|
|
||||||
)
|
)
|
||||||
|
|
||||||
@action(detail=True, methods=["get"], permission_classes=[IsAuthenticated])
|
@action(detail=True, methods=["get"], permission_classes=[IsAuthenticated])
|
||||||
def balance(self, request, pk=None):
|
def balance(self, request, pk=None):
|
||||||
"""Get current and projected balance for an account."""
|
"""Get current and projected balance for an account."""
|
||||||
account = self.get_object()
|
account = self.get_object()
|
||||||
|
|
||||||
current_balance = get_account_balance(account, paid_only=True)
|
current_balance = get_account_balance(account, paid_only=True)
|
||||||
projected_balance = get_account_balance(account, paid_only=False)
|
projected_balance = get_account_balance(account, paid_only=False)
|
||||||
|
|
||||||
serializer = AccountBalanceSerializer({
|
|
||||||
"current_balance": current_balance,
|
|
||||||
"projected_balance": projected_balance,
|
|
||||||
"currency": account.currency,
|
|
||||||
})
|
|
||||||
|
|
||||||
return Response(serializer.data)
|
|
||||||
|
|
||||||
|
serializer = AccountBalanceSerializer(
|
||||||
|
{
|
||||||
|
"current_balance": current_balance,
|
||||||
|
"projected_balance": projected_balance,
|
||||||
|
"currency": account.currency,
|
||||||
|
}
|
||||||
|
)
|
||||||
|
|
||||||
|
return Response(serializer.data)
|
||||||
|
|||||||
@@ -9,8 +9,28 @@ from apps.currencies.models import ExchangeRate
|
|||||||
class CurrencyViewSet(viewsets.ModelViewSet):
|
class CurrencyViewSet(viewsets.ModelViewSet):
|
||||||
queryset = Currency.objects.all()
|
queryset = Currency.objects.all()
|
||||||
serializer_class = CurrencySerializer
|
serializer_class = CurrencySerializer
|
||||||
|
filterset_fields = {
|
||||||
|
'name': ['exact', 'icontains'],
|
||||||
|
'code': ['exact', 'icontains'],
|
||||||
|
'decimal_places': ['exact', 'gte', 'lte', 'gt', 'lt'],
|
||||||
|
'prefix': ['exact', 'icontains'],
|
||||||
|
'suffix': ['exact', 'icontains'],
|
||||||
|
'exchange_currency': ['exact'],
|
||||||
|
'is_archived': ['exact'],
|
||||||
|
}
|
||||||
|
search_fields = '__all__'
|
||||||
|
ordering_fields = '__all__'
|
||||||
|
|
||||||
|
|
||||||
class ExchangeRateViewSet(viewsets.ModelViewSet):
|
class ExchangeRateViewSet(viewsets.ModelViewSet):
|
||||||
queryset = ExchangeRate.objects.all()
|
queryset = ExchangeRate.objects.all()
|
||||||
serializer_class = ExchangeRateSerializer
|
serializer_class = ExchangeRateSerializer
|
||||||
|
filterset_fields = {
|
||||||
|
'from_currency': ['exact'],
|
||||||
|
'to_currency': ['exact'],
|
||||||
|
'rate': ['exact', 'gte', 'lte', 'gt', 'lt'],
|
||||||
|
'date': ['exact', 'gte', 'lte', 'gt', 'lt'],
|
||||||
|
'automatic': ['exact'],
|
||||||
|
}
|
||||||
|
search_fields = '__all__'
|
||||||
|
ordering_fields = '__all__'
|
||||||
|
|||||||
@@ -2,12 +2,27 @@ from rest_framework import viewsets
|
|||||||
from rest_framework.decorators import action
|
from rest_framework.decorators import action
|
||||||
from rest_framework.response import Response
|
from rest_framework.response import Response
|
||||||
from apps.dca.models import DCAStrategy, DCAEntry
|
from apps.dca.models import DCAStrategy, DCAEntry
|
||||||
|
from apps.api.permissions import SHARED_OBJECT_PERMISSIONS
|
||||||
from apps.api.serializers import DCAStrategySerializer, DCAEntrySerializer
|
from apps.api.serializers import DCAStrategySerializer, DCAEntrySerializer
|
||||||
|
|
||||||
|
|
||||||
class DCAStrategyViewSet(viewsets.ModelViewSet):
|
class DCAStrategyViewSet(viewsets.ModelViewSet):
|
||||||
|
permission_classes = SHARED_OBJECT_PERMISSIONS
|
||||||
queryset = DCAStrategy.objects.all()
|
queryset = DCAStrategy.objects.all()
|
||||||
serializer_class = DCAStrategySerializer
|
serializer_class = DCAStrategySerializer
|
||||||
|
filterset_fields = {
|
||||||
|
"name": ["exact", "icontains"],
|
||||||
|
"target_currency": ["exact"],
|
||||||
|
"payment_currency": ["exact"],
|
||||||
|
"notes": ["exact", "icontains"],
|
||||||
|
"created_at": ["exact", "gte", "lte", "gt", "lt"],
|
||||||
|
"updated_at": ["exact", "gte", "lte", "gt", "lt"],
|
||||||
|
}
|
||||||
|
search_fields = ["name", "notes"]
|
||||||
|
ordering_fields = "__all__"
|
||||||
|
|
||||||
|
def get_queryset(self):
|
||||||
|
return DCAStrategy.objects.all()
|
||||||
|
|
||||||
@action(detail=True, methods=["get"])
|
@action(detail=True, methods=["get"])
|
||||||
def investment_frequency(self, request, pk=None):
|
def investment_frequency(self, request, pk=None):
|
||||||
@@ -30,12 +45,26 @@ class DCAStrategyViewSet(viewsets.ModelViewSet):
|
|||||||
|
|
||||||
|
|
||||||
class DCAEntryViewSet(viewsets.ModelViewSet):
|
class DCAEntryViewSet(viewsets.ModelViewSet):
|
||||||
|
permission_classes = SHARED_OBJECT_PERMISSIONS
|
||||||
|
shared_object_via = "strategy"
|
||||||
queryset = DCAEntry.objects.all()
|
queryset = DCAEntry.objects.all()
|
||||||
serializer_class = DCAEntrySerializer
|
serializer_class = DCAEntrySerializer
|
||||||
|
filterset_fields = {
|
||||||
|
"strategy": ["exact"],
|
||||||
|
"date": ["exact", "gte", "lte", "gt", "lt"],
|
||||||
|
"amount_paid": ["exact", "gte", "lte", "gt", "lt"],
|
||||||
|
"amount_received": ["exact", "gte", "lte", "gt", "lt"],
|
||||||
|
"expense_transaction": ["exact", "isnull"],
|
||||||
|
"income_transaction": ["exact", "isnull"],
|
||||||
|
"notes": ["exact", "icontains"],
|
||||||
|
"created_at": ["exact", "gte", "lte", "gt", "lt"],
|
||||||
|
"updated_at": ["exact", "gte", "lte", "gt", "lt"],
|
||||||
|
}
|
||||||
|
search_fields = ["notes"]
|
||||||
|
ordering_fields = "__all__"
|
||||||
|
ordering = ["-date"]
|
||||||
|
|
||||||
def get_queryset(self):
|
def get_queryset(self):
|
||||||
queryset = DCAEntry.objects.all()
|
# Filter entries by strategies the user has access to
|
||||||
strategy_id = self.request.query_params.get("strategy", None)
|
accessible_strategies = DCAStrategy.objects.all()
|
||||||
if strategy_id is not None:
|
return DCAEntry.objects.filter(strategy__in=accessible_strategies)
|
||||||
queryset = queryset.filter(strategy_id=strategy_id)
|
|
||||||
return queryset
|
|
||||||
|
|||||||
@@ -28,6 +28,14 @@ class ImportProfileViewSet(viewsets.ReadOnlyModelViewSet):
|
|||||||
queryset = ImportProfile.objects.all()
|
queryset = ImportProfile.objects.all()
|
||||||
serializer_class = ImportProfileSerializer
|
serializer_class = ImportProfileSerializer
|
||||||
permission_classes = [IsAuthenticated]
|
permission_classes = [IsAuthenticated]
|
||||||
|
filterset_fields = {
|
||||||
|
'name': ['exact', 'icontains'],
|
||||||
|
'yaml_config': ['exact', 'icontains'],
|
||||||
|
'version': ['exact'],
|
||||||
|
}
|
||||||
|
search_fields = ['name', 'yaml_config']
|
||||||
|
ordering_fields = '__all__'
|
||||||
|
ordering = ['name']
|
||||||
|
|
||||||
|
|
||||||
@extend_schema_view(
|
@extend_schema_view(
|
||||||
@@ -55,6 +63,22 @@ class ImportRunViewSet(viewsets.ReadOnlyModelViewSet):
|
|||||||
queryset = ImportRun.objects.all().order_by("-id")
|
queryset = ImportRun.objects.all().order_by("-id")
|
||||||
serializer_class = ImportRunSerializer
|
serializer_class = ImportRunSerializer
|
||||||
permission_classes = [IsAuthenticated]
|
permission_classes = [IsAuthenticated]
|
||||||
|
filterset_fields = {
|
||||||
|
'status': ['exact'],
|
||||||
|
'profile': ['exact'],
|
||||||
|
'file_name': ['exact', 'icontains'],
|
||||||
|
'logs': ['exact', 'icontains'],
|
||||||
|
'processed_rows': ['exact', 'gte', 'lte', 'gt', 'lt'],
|
||||||
|
'total_rows': ['exact', 'gte', 'lte', 'gt', 'lt'],
|
||||||
|
'successful_rows': ['exact', 'gte', 'lte', 'gt', 'lt'],
|
||||||
|
'skipped_rows': ['exact', 'gte', 'lte', 'gt', 'lt'],
|
||||||
|
'failed_rows': ['exact', 'gte', 'lte', 'gt', 'lt'],
|
||||||
|
'started_at': ['exact', 'gte', 'lte', 'gt', 'lt', 'isnull'],
|
||||||
|
'finished_at': ['exact', 'gte', 'lte', 'gt', 'lt', 'isnull'],
|
||||||
|
}
|
||||||
|
search_fields = ['file_name', 'logs']
|
||||||
|
ordering_fields = '__all__'
|
||||||
|
ordering = ['-id']
|
||||||
|
|
||||||
def get_queryset(self):
|
def get_queryset(self):
|
||||||
queryset = super().get_queryset()
|
queryset = super().get_queryset()
|
||||||
|
|||||||
@@ -2,7 +2,6 @@ from copy import deepcopy
|
|||||||
|
|
||||||
from rest_framework import viewsets
|
from rest_framework import viewsets
|
||||||
|
|
||||||
from apps.api.custom.pagination import CustomPageNumberPagination
|
|
||||||
from apps.api.serializers import (
|
from apps.api.serializers import (
|
||||||
TransactionSerializer,
|
TransactionSerializer,
|
||||||
TransactionCategorySerializer,
|
TransactionCategorySerializer,
|
||||||
@@ -20,12 +19,40 @@ from apps.transactions.models import (
|
|||||||
RecurringTransaction,
|
RecurringTransaction,
|
||||||
)
|
)
|
||||||
from apps.rules.signals import transaction_updated, transaction_created
|
from apps.rules.signals import transaction_updated, transaction_created
|
||||||
|
from apps.api.permissions import SHARED_OBJECT_PERMISSIONS
|
||||||
|
|
||||||
|
|
||||||
class TransactionViewSet(viewsets.ModelViewSet):
|
class TransactionViewSet(viewsets.ModelViewSet):
|
||||||
queryset = Transaction.objects.all()
|
queryset = Transaction.objects.all()
|
||||||
serializer_class = TransactionSerializer
|
serializer_class = TransactionSerializer
|
||||||
pagination_class = CustomPageNumberPagination
|
filterset_fields = {
|
||||||
|
"account": ["exact"],
|
||||||
|
"type": ["exact"],
|
||||||
|
"is_paid": ["exact"],
|
||||||
|
"date": ["exact", "gte", "lte", "gt", "lt"],
|
||||||
|
"reference_date": ["exact", "gte", "lte", "gt", "lt"],
|
||||||
|
"mute": ["exact"],
|
||||||
|
"amount": ["exact", "gte", "lte", "gt", "lt"],
|
||||||
|
"description": ["exact", "icontains"],
|
||||||
|
"notes": ["exact", "icontains"],
|
||||||
|
"category": ["exact", "isnull"],
|
||||||
|
"installment_plan": ["exact", "isnull"],
|
||||||
|
"installment_id": ["exact", "gte", "lte"],
|
||||||
|
"recurring_transaction": ["exact", "isnull"],
|
||||||
|
"internal_note": ["exact", "icontains"],
|
||||||
|
"internal_id": ["exact"],
|
||||||
|
"deleted": ["exact"],
|
||||||
|
"created_at": ["exact", "gte", "lte", "gt", "lt"],
|
||||||
|
"updated_at": ["exact", "gte", "lte", "gt", "lt"],
|
||||||
|
"deleted_at": ["exact", "gte", "lte", "gt", "lt", "isnull"],
|
||||||
|
"owner": ["exact"],
|
||||||
|
}
|
||||||
|
search_fields = ["description", "notes", "internal_note"]
|
||||||
|
ordering_fields = "__all__"
|
||||||
|
ordering = ["-id"]
|
||||||
|
|
||||||
|
def get_queryset(self):
|
||||||
|
return Transaction.objects.all()
|
||||||
|
|
||||||
def perform_create(self, serializer):
|
def perform_create(self, serializer):
|
||||||
instance = serializer.save()
|
instance = serializer.save()
|
||||||
@@ -40,50 +67,112 @@ class TransactionViewSet(viewsets.ModelViewSet):
|
|||||||
kwargs["partial"] = True
|
kwargs["partial"] = True
|
||||||
return self.update(request, *args, **kwargs)
|
return self.update(request, *args, **kwargs)
|
||||||
|
|
||||||
def get_queryset(self):
|
|
||||||
return Transaction.objects.all().order_by("-id")
|
|
||||||
|
|
||||||
|
|
||||||
class TransactionCategoryViewSet(viewsets.ModelViewSet):
|
class TransactionCategoryViewSet(viewsets.ModelViewSet):
|
||||||
|
permission_classes = SHARED_OBJECT_PERMISSIONS
|
||||||
queryset = TransactionCategory.objects.all()
|
queryset = TransactionCategory.objects.all()
|
||||||
serializer_class = TransactionCategorySerializer
|
serializer_class = TransactionCategorySerializer
|
||||||
pagination_class = CustomPageNumberPagination
|
filterset_fields = {
|
||||||
|
"name": ["exact", "icontains"],
|
||||||
|
"mute": ["exact"],
|
||||||
|
"active": ["exact"],
|
||||||
|
"owner": ["exact"],
|
||||||
|
}
|
||||||
|
search_fields = ["name"]
|
||||||
|
ordering_fields = "__all__"
|
||||||
|
ordering = ["id"]
|
||||||
|
|
||||||
def get_queryset(self):
|
def get_queryset(self):
|
||||||
return TransactionCategory.objects.all().order_by("id")
|
return TransactionCategory.objects.all()
|
||||||
|
|
||||||
|
|
||||||
class TransactionTagViewSet(viewsets.ModelViewSet):
|
class TransactionTagViewSet(viewsets.ModelViewSet):
|
||||||
|
permission_classes = SHARED_OBJECT_PERMISSIONS
|
||||||
queryset = TransactionTag.objects.all()
|
queryset = TransactionTag.objects.all()
|
||||||
serializer_class = TransactionTagSerializer
|
serializer_class = TransactionTagSerializer
|
||||||
pagination_class = CustomPageNumberPagination
|
filterset_fields = {
|
||||||
|
"name": ["exact", "icontains"],
|
||||||
|
"active": ["exact"],
|
||||||
|
"owner": ["exact"],
|
||||||
|
}
|
||||||
|
search_fields = ["name"]
|
||||||
|
ordering_fields = "__all__"
|
||||||
|
ordering = ["id"]
|
||||||
|
|
||||||
def get_queryset(self):
|
def get_queryset(self):
|
||||||
return TransactionTag.objects.all().order_by("id")
|
return TransactionTag.objects.all()
|
||||||
|
|
||||||
|
|
||||||
class TransactionEntityViewSet(viewsets.ModelViewSet):
|
class TransactionEntityViewSet(viewsets.ModelViewSet):
|
||||||
|
permission_classes = SHARED_OBJECT_PERMISSIONS
|
||||||
queryset = TransactionEntity.objects.all()
|
queryset = TransactionEntity.objects.all()
|
||||||
serializer_class = TransactionEntitySerializer
|
serializer_class = TransactionEntitySerializer
|
||||||
pagination_class = CustomPageNumberPagination
|
filterset_fields = {
|
||||||
|
"name": ["exact", "icontains"],
|
||||||
|
"active": ["exact"],
|
||||||
|
"owner": ["exact"],
|
||||||
|
}
|
||||||
|
search_fields = ["name"]
|
||||||
|
ordering_fields = "__all__"
|
||||||
|
ordering = ["id"]
|
||||||
|
|
||||||
def get_queryset(self):
|
def get_queryset(self):
|
||||||
return TransactionEntity.objects.all().order_by("id")
|
return TransactionEntity.objects.all()
|
||||||
|
|
||||||
|
|
||||||
class InstallmentPlanViewSet(viewsets.ModelViewSet):
|
class InstallmentPlanViewSet(viewsets.ModelViewSet):
|
||||||
queryset = InstallmentPlan.objects.all()
|
queryset = InstallmentPlan.objects.all()
|
||||||
serializer_class = InstallmentPlanSerializer
|
serializer_class = InstallmentPlanSerializer
|
||||||
pagination_class = CustomPageNumberPagination
|
filterset_fields = {
|
||||||
|
"account": ["exact"],
|
||||||
|
"type": ["exact"],
|
||||||
|
"description": ["exact", "icontains"],
|
||||||
|
"number_of_installments": ["exact", "gte", "lte", "gt", "lt"],
|
||||||
|
"installment_start": ["exact", "gte", "lte", "gt", "lt"],
|
||||||
|
"installment_total_number": ["exact", "gte", "lte", "gt", "lt"],
|
||||||
|
"start_date": ["exact", "gte", "lte", "gt", "lt"],
|
||||||
|
"reference_date": ["exact", "gte", "lte", "gt", "lt", "isnull"],
|
||||||
|
"end_date": ["exact", "gte", "lte", "gt", "lt", "isnull"],
|
||||||
|
"recurrence": ["exact"],
|
||||||
|
"installment_amount": ["exact", "gte", "lte", "gt", "lt"],
|
||||||
|
"category": ["exact", "isnull"],
|
||||||
|
"notes": ["exact", "icontains"],
|
||||||
|
"add_description_to_transaction": ["exact"],
|
||||||
|
"add_notes_to_transaction": ["exact"],
|
||||||
|
}
|
||||||
|
search_fields = ["description", "notes"]
|
||||||
|
ordering_fields = "__all__"
|
||||||
|
ordering = ["-id"]
|
||||||
|
|
||||||
def get_queryset(self):
|
def get_queryset(self):
|
||||||
return InstallmentPlan.objects.all().order_by("-id")
|
return InstallmentPlan.objects.all()
|
||||||
|
|
||||||
|
|
||||||
class RecurringTransactionViewSet(viewsets.ModelViewSet):
|
class RecurringTransactionViewSet(viewsets.ModelViewSet):
|
||||||
queryset = RecurringTransaction.objects.all()
|
queryset = RecurringTransaction.objects.all()
|
||||||
serializer_class = RecurringTransactionSerializer
|
serializer_class = RecurringTransactionSerializer
|
||||||
pagination_class = CustomPageNumberPagination
|
filterset_fields = {
|
||||||
|
"is_paused": ["exact"],
|
||||||
|
"account": ["exact"],
|
||||||
|
"type": ["exact"],
|
||||||
|
"amount": ["exact", "gte", "lte", "gt", "lt"],
|
||||||
|
"description": ["exact", "icontains"],
|
||||||
|
"category": ["exact", "isnull"],
|
||||||
|
"notes": ["exact", "icontains"],
|
||||||
|
"reference_date": ["exact", "gte", "lte", "gt", "lt", "isnull"],
|
||||||
|
"start_date": ["exact", "gte", "lte", "gt", "lt"],
|
||||||
|
"end_date": ["exact", "gte", "lte", "gt", "lt", "isnull"],
|
||||||
|
"recurrence_type": ["exact"],
|
||||||
|
"recurrence_interval": ["exact", "gte", "lte", "gt", "lt"],
|
||||||
|
"keep_at_most": ["exact", "gte", "lte", "gt", "lt"],
|
||||||
|
"last_generated_date": ["exact", "gte", "lte", "gt", "lt", "isnull"],
|
||||||
|
"last_generated_reference_date": ["exact", "gte", "lte", "gt", "lt", "isnull"],
|
||||||
|
"add_description_to_transaction": ["exact"],
|
||||||
|
"add_notes_to_transaction": ["exact"],
|
||||||
|
}
|
||||||
|
search_fields = ["description", "notes"]
|
||||||
|
ordering_fields = "__all__"
|
||||||
|
ordering = ["-id"]
|
||||||
|
|
||||||
def get_queryset(self):
|
def get_queryset(self):
|
||||||
return RecurringTransaction.objects.all().order_by("-id")
|
return RecurringTransaction.objects.all()
|
||||||
|
|||||||
@@ -23,3 +23,6 @@ class CommonConfig(AppConfig):
|
|||||||
# Delete the cache for update checks to prevent false-positives when the app is restarted
|
# Delete the cache for update checks to prevent false-positives when the app is restarted
|
||||||
# this will be recreated by the check_for_updates task
|
# this will be recreated by the check_for_updates task
|
||||||
cache.delete("update_check")
|
cache.delete("update_check")
|
||||||
|
|
||||||
|
# Register system checks for required environment variables
|
||||||
|
from apps.common import checks # noqa: F401
|
||||||
|
|||||||
@@ -0,0 +1,103 @@
|
|||||||
|
"""
|
||||||
|
Django System Checks for required environment variables.
|
||||||
|
|
||||||
|
This module validates that required environment variables (those without defaults)
|
||||||
|
are present before the application starts.
|
||||||
|
"""
|
||||||
|
|
||||||
|
import os
|
||||||
|
|
||||||
|
from django.core.checks import Error, register
|
||||||
|
|
||||||
|
|
||||||
|
# List of environment variables that are required (no default values)
|
||||||
|
# Based on the README.md documentation
|
||||||
|
REQUIRED_ENV_VARS = [
|
||||||
|
("SECRET_KEY", "This is used to provide cryptographic signing."),
|
||||||
|
("SQL_DATABASE", "The name of your postgres database."),
|
||||||
|
]
|
||||||
|
|
||||||
|
# List of environment variables that must be valid integers if set
|
||||||
|
INT_ENV_VARS = [
|
||||||
|
("TASK_WORKERS", "How many workers to have for async tasks."),
|
||||||
|
("SESSION_EXPIRY_TIME", "The age of session cookies, in seconds."),
|
||||||
|
("INTERNAL_PORT", "The port on which the app listens on."),
|
||||||
|
("DJANGO_VITE_DEV_SERVER_PORT", "The port where Vite's dev server is running"),
|
||||||
|
]
|
||||||
|
|
||||||
|
|
||||||
|
@register()
|
||||||
|
def check_required_env_vars(app_configs, **kwargs):
|
||||||
|
"""
|
||||||
|
Check that all required environment variables are set.
|
||||||
|
|
||||||
|
Returns a list of Error objects for any missing required variables.
|
||||||
|
"""
|
||||||
|
errors = []
|
||||||
|
|
||||||
|
for var_name, description in REQUIRED_ENV_VARS:
|
||||||
|
value = os.getenv(var_name)
|
||||||
|
if not value:
|
||||||
|
errors.append(
|
||||||
|
Error(
|
||||||
|
f"Required environment variable '{var_name}' is not set.",
|
||||||
|
hint=f"{description} Please set this variable in your .env file or environment.",
|
||||||
|
id="wygiwyh.E001",
|
||||||
|
)
|
||||||
|
)
|
||||||
|
|
||||||
|
return errors
|
||||||
|
|
||||||
|
|
||||||
|
@register()
|
||||||
|
def check_int_env_vars(app_configs, **kwargs):
|
||||||
|
"""
|
||||||
|
Check that environment variables that should be integers are valid.
|
||||||
|
|
||||||
|
Returns a list of Error objects for any invalid integer variables.
|
||||||
|
"""
|
||||||
|
errors = []
|
||||||
|
|
||||||
|
for var_name, description in INT_ENV_VARS:
|
||||||
|
value = os.getenv(var_name)
|
||||||
|
if value is not None:
|
||||||
|
try:
|
||||||
|
int(value)
|
||||||
|
except ValueError:
|
||||||
|
errors.append(
|
||||||
|
Error(
|
||||||
|
f"Environment variable '{var_name}' must be a valid integer, got '{value}'.",
|
||||||
|
hint=f"{description}",
|
||||||
|
id="wygiwyh.E002",
|
||||||
|
)
|
||||||
|
)
|
||||||
|
|
||||||
|
return errors
|
||||||
|
|
||||||
|
|
||||||
|
@register()
|
||||||
|
def check_soft_delete_config(app_configs, **kwargs):
|
||||||
|
"""
|
||||||
|
Check that KEEP_DELETED_TRANSACTIONS_FOR is a valid integer when ENABLE_SOFT_DELETE is enabled.
|
||||||
|
|
||||||
|
Returns a list of Error objects if the configuration is invalid.
|
||||||
|
"""
|
||||||
|
errors = []
|
||||||
|
|
||||||
|
enable_soft_delete = os.getenv("ENABLE_SOFT_DELETE", "false").lower() == "true"
|
||||||
|
|
||||||
|
if enable_soft_delete:
|
||||||
|
keep_deleted_for = os.getenv("KEEP_DELETED_TRANSACTIONS_FOR")
|
||||||
|
if keep_deleted_for is not None:
|
||||||
|
try:
|
||||||
|
int(keep_deleted_for)
|
||||||
|
except ValueError:
|
||||||
|
errors.append(
|
||||||
|
Error(
|
||||||
|
f"Environment variable 'KEEP_DELETED_TRANSACTIONS_FOR' must be a valid integer when ENABLE_SOFT_DELETE is enabled, got '{keep_deleted_for}'.",
|
||||||
|
hint="Time in days to keep soft deleted transactions for. Set to 0 to keep all transactions indefinitely.",
|
||||||
|
id="wygiwyh.E003",
|
||||||
|
)
|
||||||
|
)
|
||||||
|
|
||||||
|
return errors
|
||||||
@@ -0,0 +1,58 @@
|
|||||||
|
from django.core.exceptions import PermissionDenied
|
||||||
|
from django.http import Http404
|
||||||
|
from django.shortcuts import get_object_or_404
|
||||||
|
|
||||||
|
READ = "read"
|
||||||
|
EDIT = "edit"
|
||||||
|
|
||||||
|
|
||||||
|
def get_shared_object_or_error(klass, request, *, level=EDIT, via=None, **kwargs):
|
||||||
|
"""Fetch an object like ``get_object_or_404`` while enforcing access control.
|
||||||
|
|
||||||
|
``SharedObjectManager`` scopes querysets to what a user may *see*, which is
|
||||||
|
not the same as what they may *change*. Views that resolve an object from a
|
||||||
|
URL id must state which of the two they need, otherwise a shared or public
|
||||||
|
object becomes writable by anyone who can see it.
|
||||||
|
|
||||||
|
``level`` selects the check applied to the governing ``SharedObject``:
|
||||||
|
|
||||||
|
``READ``
|
||||||
|
The object must be visible to the user. Denial raises :class:`Http404`
|
||||||
|
so the response does not confirm that the id exists.
|
||||||
|
``EDIT``
|
||||||
|
The object must be owned by the user. Denial raises
|
||||||
|
:class:`~django.core.exceptions.PermissionDenied` (HTTP 403), which the
|
||||||
|
frontend surfaces as an "Access Denied" dialog. An object the user
|
||||||
|
cannot even see raises :class:`Http404` instead, so 403 never confirms
|
||||||
|
the existence of an object they were not allowed to know about.
|
||||||
|
|
||||||
|
Objects with no owner stay accessible to everyone, preserving the existing
|
||||||
|
behaviour for legacy/unowned objects.
|
||||||
|
|
||||||
|
``via`` is a dotted path to the ``SharedObject`` that governs access, for
|
||||||
|
models owned through a relation, e.g. a rule action governed by its parent
|
||||||
|
rule::
|
||||||
|
|
||||||
|
get_shared_object_or_error(
|
||||||
|
TransactionRuleAction, request, id=pk, level=EDIT, via="rule"
|
||||||
|
)
|
||||||
|
|
||||||
|
The path is resolved with a plain ``getattr``, so a path that does not
|
||||||
|
resolve raises ``AttributeError`` rather than silently granting access.
|
||||||
|
"""
|
||||||
|
obj = get_object_or_404(klass, **kwargs)
|
||||||
|
|
||||||
|
guard = obj
|
||||||
|
for attr in via.split(".") if via else []:
|
||||||
|
guard = getattr(guard, attr)
|
||||||
|
|
||||||
|
if level not in (READ, EDIT):
|
||||||
|
raise ValueError(f"Unknown access level: {level!r}")
|
||||||
|
|
||||||
|
if not guard.is_visible_to(request.user):
|
||||||
|
raise Http404
|
||||||
|
|
||||||
|
if level == EDIT and not guard.is_editable_by(request.user):
|
||||||
|
raise PermissionDenied
|
||||||
|
|
||||||
|
return obj
|
||||||
@@ -0,0 +1,58 @@
|
|||||||
|
from datetime import timedelta
|
||||||
|
|
||||||
|
from django.contrib.auth import get_user_model
|
||||||
|
from django.core.management.base import BaseCommand, CommandError
|
||||||
|
from django.utils import timezone
|
||||||
|
|
||||||
|
from apps.users.models import APIToken
|
||||||
|
|
||||||
|
|
||||||
|
class Command(BaseCommand):
|
||||||
|
help = "Creates a hashed API token for a WYGIWYH user and prints the raw token once."
|
||||||
|
|
||||||
|
def add_arguments(self, parser):
|
||||||
|
parser.add_argument("email", help="WYGIWYH user email that will own this token.")
|
||||||
|
parser.add_argument(
|
||||||
|
"--name",
|
||||||
|
default="n8n",
|
||||||
|
help="Human-readable token name. Defaults to 'n8n'.",
|
||||||
|
)
|
||||||
|
parser.add_argument(
|
||||||
|
"--expires-in-days",
|
||||||
|
type=int,
|
||||||
|
default=None,
|
||||||
|
help="Optional token lifetime in whole days.",
|
||||||
|
)
|
||||||
|
|
||||||
|
def handle(self, *args, **options):
|
||||||
|
email = options["email"].strip()
|
||||||
|
name = options["name"].strip()
|
||||||
|
expires_in_days = options["expires_in_days"]
|
||||||
|
|
||||||
|
if not email:
|
||||||
|
raise CommandError("Email is required.")
|
||||||
|
if not name:
|
||||||
|
raise CommandError("Token name cannot be empty.")
|
||||||
|
if expires_in_days is not None and expires_in_days <= 0:
|
||||||
|
raise CommandError("--expires-in-days must be greater than zero.")
|
||||||
|
|
||||||
|
user = get_user_model().objects.filter(email__iexact=email).first()
|
||||||
|
if user is None:
|
||||||
|
raise CommandError(f"No WYGIWYH user exists for '{email}'.")
|
||||||
|
|
||||||
|
expires_at = None
|
||||||
|
if expires_in_days is not None:
|
||||||
|
expires_at = timezone.now() + timedelta(days=expires_in_days)
|
||||||
|
|
||||||
|
token, raw_token = APIToken.objects.create_token(
|
||||||
|
user=user,
|
||||||
|
name=name,
|
||||||
|
expires_at=expires_at,
|
||||||
|
)
|
||||||
|
|
||||||
|
self.stdout.write(
|
||||||
|
self.style.SUCCESS(
|
||||||
|
f"Created API token '{token.name}' for {user.email} ({token.token_key})."
|
||||||
|
)
|
||||||
|
)
|
||||||
|
self.stdout.write(raw_token)
|
||||||
@@ -0,0 +1,129 @@
|
|||||||
|
import os
|
||||||
|
|
||||||
|
from django.contrib.auth.hashers import check_password
|
||||||
|
from django.core.exceptions import ValidationError
|
||||||
|
from django.core.management.base import BaseCommand, CommandError
|
||||||
|
from oauth2_provider.models import get_application_model
|
||||||
|
|
||||||
|
|
||||||
|
Application = get_application_model()
|
||||||
|
|
||||||
|
|
||||||
|
def _get_env(name: str) -> str:
|
||||||
|
return os.getenv(name, "").strip()
|
||||||
|
|
||||||
|
|
||||||
|
def _get_bool_env(name: str, default: bool = False) -> bool:
|
||||||
|
raw = _get_env(name)
|
||||||
|
if not raw:
|
||||||
|
return default
|
||||||
|
return raw.lower() in {"1", "true", "yes", "on"}
|
||||||
|
|
||||||
|
|
||||||
|
class Command(BaseCommand):
|
||||||
|
help = (
|
||||||
|
"Creates or updates the OAuth application used by MCP clients when "
|
||||||
|
"MCP_OAUTH_CLIENT_* environment variables are configured."
|
||||||
|
)
|
||||||
|
|
||||||
|
def handle(self, *args, **options):
|
||||||
|
client_id = _get_env("MCP_OAUTH_CLIENT_ID")
|
||||||
|
client_secret = _get_env("MCP_OAUTH_CLIENT_SECRET")
|
||||||
|
redirect_uris = " ".join(_get_env("MCP_OAUTH_REDIRECT_URIS").split())
|
||||||
|
name = _get_env("MCP_OAUTH_CLIENT_NAME") or "WYGIWYH MCP"
|
||||||
|
skip_authorization = _get_bool_env("MCP_OAUTH_SKIP_AUTHORIZATION", default=False)
|
||||||
|
|
||||||
|
if not any([client_id, client_secret, redirect_uris]):
|
||||||
|
self.stdout.write(
|
||||||
|
self.style.NOTICE(
|
||||||
|
"MCP OAuth client env vars are not set. Skipping OAuth application setup."
|
||||||
|
)
|
||||||
|
)
|
||||||
|
return
|
||||||
|
|
||||||
|
missing = []
|
||||||
|
if not client_id:
|
||||||
|
missing.append("MCP_OAUTH_CLIENT_ID")
|
||||||
|
if not client_secret:
|
||||||
|
missing.append("MCP_OAUTH_CLIENT_SECRET")
|
||||||
|
if not redirect_uris:
|
||||||
|
missing.append("MCP_OAUTH_REDIRECT_URIS")
|
||||||
|
if missing:
|
||||||
|
raise CommandError(
|
||||||
|
"Missing required MCP OAuth settings: " + ", ".join(missing)
|
||||||
|
)
|
||||||
|
|
||||||
|
application, created = Application.objects.get_or_create(
|
||||||
|
client_id=client_id,
|
||||||
|
defaults={
|
||||||
|
"name": name,
|
||||||
|
"client_type": Application.CLIENT_CONFIDENTIAL,
|
||||||
|
"authorization_grant_type": Application.GRANT_AUTHORIZATION_CODE,
|
||||||
|
"redirect_uris": redirect_uris,
|
||||||
|
"skip_authorization": skip_authorization,
|
||||||
|
"client_secret": client_secret,
|
||||||
|
"hash_client_secret": True,
|
||||||
|
},
|
||||||
|
)
|
||||||
|
|
||||||
|
updated_fields = []
|
||||||
|
if application.name != name:
|
||||||
|
application.name = name
|
||||||
|
updated_fields.append("name")
|
||||||
|
if application.client_type != Application.CLIENT_CONFIDENTIAL:
|
||||||
|
application.client_type = Application.CLIENT_CONFIDENTIAL
|
||||||
|
updated_fields.append("client_type")
|
||||||
|
if (
|
||||||
|
application.authorization_grant_type
|
||||||
|
!= Application.GRANT_AUTHORIZATION_CODE
|
||||||
|
):
|
||||||
|
application.authorization_grant_type = Application.GRANT_AUTHORIZATION_CODE
|
||||||
|
updated_fields.append("authorization_grant_type")
|
||||||
|
if application.redirect_uris != redirect_uris:
|
||||||
|
application.redirect_uris = redirect_uris
|
||||||
|
updated_fields.append("redirect_uris")
|
||||||
|
if application.skip_authorization != skip_authorization:
|
||||||
|
application.skip_authorization = skip_authorization
|
||||||
|
updated_fields.append("skip_authorization")
|
||||||
|
if application.hash_client_secret is not True:
|
||||||
|
application.hash_client_secret = True
|
||||||
|
updated_fields.append("hash_client_secret")
|
||||||
|
if not application.client_secret or not check_password(
|
||||||
|
client_secret,
|
||||||
|
application.client_secret,
|
||||||
|
):
|
||||||
|
application.client_secret = client_secret
|
||||||
|
updated_fields.append("client_secret")
|
||||||
|
|
||||||
|
try:
|
||||||
|
application.full_clean()
|
||||||
|
except ValidationError as exc:
|
||||||
|
errors = "; ".join(
|
||||||
|
f"{field}: {', '.join(messages)}"
|
||||||
|
for field, messages in exc.message_dict.items()
|
||||||
|
)
|
||||||
|
raise CommandError(f"Invalid MCP OAuth application settings: {errors}") from exc
|
||||||
|
|
||||||
|
if created:
|
||||||
|
application.save()
|
||||||
|
self.stdout.write(
|
||||||
|
self.style.SUCCESS(
|
||||||
|
f"Created MCP OAuth application '{application.client_id}'."
|
||||||
|
)
|
||||||
|
)
|
||||||
|
return
|
||||||
|
|
||||||
|
if updated_fields:
|
||||||
|
application.save(update_fields=updated_fields)
|
||||||
|
self.stdout.write(
|
||||||
|
self.style.SUCCESS(
|
||||||
|
f"Updated MCP OAuth application '{application.client_id}'."
|
||||||
|
)
|
||||||
|
)
|
||||||
|
return
|
||||||
|
|
||||||
|
self.stdout.write(
|
||||||
|
self.style.SUCCESS(
|
||||||
|
f"MCP OAuth application '{application.client_id}' is already up to date."
|
||||||
|
)
|
||||||
|
)
|
||||||
@@ -58,13 +58,35 @@ class SharedObject(models.Model):
|
|||||||
models.Index(fields=["visibility"]),
|
models.Index(fields=["visibility"]),
|
||||||
]
|
]
|
||||||
|
|
||||||
def is_accessible_by(self, user):
|
# NOTE: these two predicates must stay in sync with the ``Q`` objects built
|
||||||
"""Check if a user can access this object"""
|
# by ``SharedObjectManager.get_queryset`` above. The manager filters at the
|
||||||
return (
|
# queryset level and these check a single instance, so they cannot share an
|
||||||
self.visibility == "public"
|
# implementation; ``SharedObjectPredicateParityTests`` asserts they agree.
|
||||||
or self.owner == user
|
def is_visible_to(self, user):
|
||||||
or (self.visibility == "shared" and user in self.shared_with.all())
|
"""Whether ``user`` may read this object.
|
||||||
)
|
|
||||||
|
Mirrors ``SharedObjectManager``: public objects, objects with no owner,
|
||||||
|
the owner's own objects, and objects explicitly shared with the user.
|
||||||
|
"""
|
||||||
|
if self.owner is None or self.visibility == "public":
|
||||||
|
return True
|
||||||
|
|
||||||
|
if not user or not user.is_authenticated:
|
||||||
|
return False
|
||||||
|
|
||||||
|
return self.owner_id == user.pk or self.shared_with.filter(pk=user.pk).exists()
|
||||||
|
|
||||||
|
def is_editable_by(self, user):
|
||||||
|
"""Whether ``user`` may mutate this object.
|
||||||
|
|
||||||
|
Sharing grants read access only; mutation stays with the owner. Objects
|
||||||
|
with no owner remain editable by everyone, preserving the behaviour of
|
||||||
|
legacy/unowned objects.
|
||||||
|
"""
|
||||||
|
if self.owner is None:
|
||||||
|
return True
|
||||||
|
|
||||||
|
return bool(user and user.is_authenticated and self.owner_id == user.pk)
|
||||||
|
|
||||||
def save(self, *args, **kwargs):
|
def save(self, *args, **kwargs):
|
||||||
if not self.pk and not self.owner:
|
if not self.pk and not self.owner:
|
||||||
|
|||||||
@@ -0,0 +1,253 @@
|
|||||||
|
import hmac
|
||||||
|
import json
|
||||||
|
import time
|
||||||
|
from secrets import token_urlsafe
|
||||||
|
|
||||||
|
from django.conf import settings
|
||||||
|
from django.core.exceptions import ValidationError
|
||||||
|
from django.http import JsonResponse
|
||||||
|
from django.views.decorators.csrf import csrf_exempt
|
||||||
|
from django.views.decorators.http import require_http_methods
|
||||||
|
from oauth2_provider.models import get_application_model
|
||||||
|
|
||||||
|
|
||||||
|
Application = get_application_model()
|
||||||
|
|
||||||
|
SUPPORTED_TOKEN_ENDPOINT_AUTH_METHODS = {
|
||||||
|
"none": Application.CLIENT_PUBLIC,
|
||||||
|
"client_secret_basic": Application.CLIENT_CONFIDENTIAL,
|
||||||
|
"client_secret_post": Application.CLIENT_CONFIDENTIAL,
|
||||||
|
}
|
||||||
|
SUPPORTED_GRANT_TYPES = {"authorization_code", "refresh_token"}
|
||||||
|
SUPPORTED_RESPONSE_TYPES = {"code"}
|
||||||
|
|
||||||
|
|
||||||
|
def _base_url(request):
|
||||||
|
return settings.PUBLIC_BASE_URL or request.build_absolute_uri("/").rstrip("/")
|
||||||
|
|
||||||
|
|
||||||
|
def _json_error(error, error_description, status=400):
|
||||||
|
response = JsonResponse(
|
||||||
|
{"error": error, "error_description": error_description},
|
||||||
|
status=status,
|
||||||
|
)
|
||||||
|
response["Cache-Control"] = "no-store"
|
||||||
|
response["Pragma"] = "no-cache"
|
||||||
|
return response
|
||||||
|
|
||||||
|
|
||||||
|
def _set_no_store_headers(response):
|
||||||
|
response["Cache-Control"] = "no-store"
|
||||||
|
response["Pragma"] = "no-cache"
|
||||||
|
return response
|
||||||
|
|
||||||
|
|
||||||
|
def _parse_json_request_body(request):
|
||||||
|
try:
|
||||||
|
payload = json.loads(request.body.decode("utf-8"))
|
||||||
|
except (UnicodeDecodeError, json.JSONDecodeError) as exc:
|
||||||
|
raise ValueError("Request body must be valid JSON.") from exc
|
||||||
|
|
||||||
|
if not isinstance(payload, dict):
|
||||||
|
raise ValueError("Request body must be a JSON object.")
|
||||||
|
|
||||||
|
return payload
|
||||||
|
|
||||||
|
|
||||||
|
def _get_string_list(payload, field_name, *, required=False, default=None):
|
||||||
|
value = payload.get(field_name, default)
|
||||||
|
if value is None:
|
||||||
|
if required:
|
||||||
|
raise ValueError(f"'{field_name}' is required.")
|
||||||
|
return None
|
||||||
|
|
||||||
|
if not isinstance(value, list) or not value:
|
||||||
|
raise ValueError(f"'{field_name}' must be a non-empty array of strings.")
|
||||||
|
|
||||||
|
normalized = []
|
||||||
|
for item in value:
|
||||||
|
if not isinstance(item, str) or not item.strip():
|
||||||
|
raise ValueError(f"'{field_name}' must contain only non-empty strings.")
|
||||||
|
normalized.append(item.strip())
|
||||||
|
return normalized
|
||||||
|
|
||||||
|
|
||||||
|
def _get_supported_scopes():
|
||||||
|
return set(settings.OAUTH2_PROVIDER.get("SCOPES", {}).keys())
|
||||||
|
|
||||||
|
|
||||||
|
def _dcr_initial_access_token_ok(request):
|
||||||
|
"""Validate the optional RFC 7591 initial access token, if one is configured."""
|
||||||
|
expected = settings.OAUTH2_DCR_INITIAL_ACCESS_TOKEN
|
||||||
|
if not expected:
|
||||||
|
return True
|
||||||
|
|
||||||
|
header = request.META.get("HTTP_AUTHORIZATION", "")
|
||||||
|
scheme, _, value = header.partition(" ")
|
||||||
|
if scheme.lower() != "bearer" or not value:
|
||||||
|
return False
|
||||||
|
return hmac.compare_digest(value, expected)
|
||||||
|
|
||||||
|
|
||||||
|
@require_http_methods(["GET"])
|
||||||
|
def authorization_server_metadata(request):
|
||||||
|
base_url = _base_url(request)
|
||||||
|
metadata = {
|
||||||
|
"issuer": base_url,
|
||||||
|
"authorization_endpoint": f"{base_url}/oauth/authorize/",
|
||||||
|
"token_endpoint": f"{base_url}/oauth/token/",
|
||||||
|
"revocation_endpoint": f"{base_url}/oauth/revoke_token/",
|
||||||
|
"introspection_endpoint": f"{base_url}/oauth/introspect/",
|
||||||
|
"scopes_supported": sorted(settings.OAUTH2_PROVIDER["SCOPES"].keys()),
|
||||||
|
"response_types_supported": ["code"],
|
||||||
|
"grant_types_supported": ["authorization_code", "refresh_token"],
|
||||||
|
"token_endpoint_auth_methods_supported": [
|
||||||
|
"none",
|
||||||
|
"client_secret_basic",
|
||||||
|
"client_secret_post",
|
||||||
|
],
|
||||||
|
"code_challenge_methods_supported": ["S256"],
|
||||||
|
}
|
||||||
|
# Only advertise registration when DCR is actually enabled.
|
||||||
|
if settings.OAUTH2_DCR_ENABLED:
|
||||||
|
metadata["registration_endpoint"] = f"{base_url}/oauth/register/"
|
||||||
|
return JsonResponse(metadata)
|
||||||
|
|
||||||
|
|
||||||
|
@csrf_exempt
|
||||||
|
@require_http_methods(["POST"])
|
||||||
|
def dynamic_client_registration(request):
|
||||||
|
if not settings.OAUTH2_DCR_ENABLED:
|
||||||
|
return _json_error(
|
||||||
|
"not_found",
|
||||||
|
"Dynamic client registration is disabled.",
|
||||||
|
status=404,
|
||||||
|
)
|
||||||
|
|
||||||
|
if not _dcr_initial_access_token_ok(request):
|
||||||
|
return _json_error(
|
||||||
|
"invalid_token",
|
||||||
|
"A valid initial access token is required to register a client.",
|
||||||
|
status=401,
|
||||||
|
)
|
||||||
|
|
||||||
|
try:
|
||||||
|
payload = _parse_json_request_body(request)
|
||||||
|
redirect_uris = _get_string_list(payload, "redirect_uris", required=True)
|
||||||
|
grant_types = _get_string_list(
|
||||||
|
payload,
|
||||||
|
"grant_types",
|
||||||
|
default=["authorization_code"],
|
||||||
|
)
|
||||||
|
response_types = _get_string_list(
|
||||||
|
payload,
|
||||||
|
"response_types",
|
||||||
|
default=["code"],
|
||||||
|
)
|
||||||
|
except ValueError as exc:
|
||||||
|
return _json_error("invalid_client_metadata", str(exc))
|
||||||
|
|
||||||
|
unsupported_grant_types = sorted(set(grant_types) - SUPPORTED_GRANT_TYPES)
|
||||||
|
if unsupported_grant_types:
|
||||||
|
return _json_error(
|
||||||
|
"invalid_client_metadata",
|
||||||
|
"Unsupported grant_types: " + ", ".join(unsupported_grant_types),
|
||||||
|
)
|
||||||
|
|
||||||
|
if "authorization_code" not in grant_types:
|
||||||
|
return _json_error(
|
||||||
|
"invalid_client_metadata",
|
||||||
|
"grant_types must include 'authorization_code'.",
|
||||||
|
)
|
||||||
|
|
||||||
|
unsupported_response_types = sorted(set(response_types) - SUPPORTED_RESPONSE_TYPES)
|
||||||
|
if unsupported_response_types:
|
||||||
|
return _json_error(
|
||||||
|
"invalid_client_metadata",
|
||||||
|
"Unsupported response_types: "
|
||||||
|
+ ", ".join(unsupported_response_types),
|
||||||
|
)
|
||||||
|
|
||||||
|
if "code" not in response_types:
|
||||||
|
return _json_error(
|
||||||
|
"invalid_client_metadata",
|
||||||
|
"response_types must include 'code'.",
|
||||||
|
)
|
||||||
|
|
||||||
|
token_endpoint_auth_method = payload.get(
|
||||||
|
"token_endpoint_auth_method",
|
||||||
|
"client_secret_basic",
|
||||||
|
)
|
||||||
|
if token_endpoint_auth_method not in SUPPORTED_TOKEN_ENDPOINT_AUTH_METHODS:
|
||||||
|
return _json_error(
|
||||||
|
"invalid_client_metadata",
|
||||||
|
"Unsupported token_endpoint_auth_method: "
|
||||||
|
+ token_endpoint_auth_method,
|
||||||
|
)
|
||||||
|
|
||||||
|
supported_scopes = _get_supported_scopes()
|
||||||
|
raw_scope = payload.get("scope", "mcp")
|
||||||
|
if not isinstance(raw_scope, str):
|
||||||
|
return _json_error(
|
||||||
|
"invalid_client_metadata",
|
||||||
|
"'scope' must be a space-delimited string.",
|
||||||
|
)
|
||||||
|
requested_scope = raw_scope.strip() or "mcp"
|
||||||
|
requested_scopes = set(requested_scope.split())
|
||||||
|
unsupported_scopes = sorted(requested_scopes - supported_scopes)
|
||||||
|
if unsupported_scopes:
|
||||||
|
return _json_error(
|
||||||
|
"invalid_client_metadata",
|
||||||
|
"Unsupported scope values: " + ", ".join(unsupported_scopes),
|
||||||
|
)
|
||||||
|
|
||||||
|
client_name = str(payload.get("client_name", "Dynamic MCP Client")).strip()
|
||||||
|
if not client_name:
|
||||||
|
client_name = "Dynamic MCP Client"
|
||||||
|
|
||||||
|
client_secret = None
|
||||||
|
client_type = SUPPORTED_TOKEN_ENDPOINT_AUTH_METHODS[token_endpoint_auth_method]
|
||||||
|
if client_type == Application.CLIENT_CONFIDENTIAL:
|
||||||
|
client_secret = token_urlsafe(48)
|
||||||
|
|
||||||
|
application = Application(
|
||||||
|
name=client_name,
|
||||||
|
client_type=client_type,
|
||||||
|
authorization_grant_type=Application.GRANT_AUTHORIZATION_CODE,
|
||||||
|
redirect_uris=" ".join(redirect_uris),
|
||||||
|
skip_authorization=False,
|
||||||
|
hash_client_secret=True,
|
||||||
|
client_secret=client_secret or "",
|
||||||
|
)
|
||||||
|
|
||||||
|
try:
|
||||||
|
application.full_clean()
|
||||||
|
except ValidationError as exc:
|
||||||
|
errors = []
|
||||||
|
for field, messages in exc.message_dict.items():
|
||||||
|
errors.extend(f"{field}: {message}" for message in messages)
|
||||||
|
return _json_error(
|
||||||
|
"invalid_client_metadata",
|
||||||
|
"; ".join(errors),
|
||||||
|
)
|
||||||
|
|
||||||
|
application.save()
|
||||||
|
|
||||||
|
response_payload = {
|
||||||
|
"client_id": application.client_id,
|
||||||
|
"client_id_issued_at": int(time.time()),
|
||||||
|
"client_name": client_name,
|
||||||
|
"redirect_uris": redirect_uris,
|
||||||
|
# Report what was actually provisioned, not the raw request echo. The app
|
||||||
|
# is created with the authorization_code grant; refresh_token is implicit
|
||||||
|
# to that grant in django-oauth-toolkit rather than a separate capability.
|
||||||
|
"grant_types": sorted(set(grant_types) & SUPPORTED_GRANT_TYPES),
|
||||||
|
"response_types": sorted(set(response_types) & SUPPORTED_RESPONSE_TYPES),
|
||||||
|
"scope": " ".join(sorted(requested_scopes)),
|
||||||
|
"token_endpoint_auth_method": token_endpoint_auth_method,
|
||||||
|
}
|
||||||
|
if client_secret is not None:
|
||||||
|
response_payload["client_secret"] = client_secret
|
||||||
|
response_payload["client_secret_expires_at"] = 0
|
||||||
|
|
||||||
|
return _set_no_store_headers(JsonResponse(response_payload, status=201))
|
||||||
@@ -1,6 +1,47 @@
|
|||||||
|
import functools
|
||||||
|
import inspect
|
||||||
|
|
||||||
import procrastinate
|
import procrastinate
|
||||||
|
from django.db import close_old_connections
|
||||||
|
|
||||||
|
|
||||||
|
_CONNECTION_CLEANUP_WRAPPED = "_wygiwyh_connection_cleanup_wrapped"
|
||||||
|
|
||||||
|
|
||||||
|
def _wrap_task_with_django_connection_cleanup(task):
|
||||||
|
if getattr(task.func, _CONNECTION_CLEANUP_WRAPPED, False):
|
||||||
|
return
|
||||||
|
|
||||||
|
func = task.func
|
||||||
|
|
||||||
|
if inspect.iscoroutinefunction(func):
|
||||||
|
|
||||||
|
@functools.wraps(func)
|
||||||
|
async def async_wrapped(*args, **kwargs):
|
||||||
|
close_old_connections()
|
||||||
|
try:
|
||||||
|
return await func(*args, **kwargs)
|
||||||
|
finally:
|
||||||
|
close_old_connections()
|
||||||
|
|
||||||
|
wrapped = async_wrapped
|
||||||
|
else:
|
||||||
|
|
||||||
|
@functools.wraps(func)
|
||||||
|
def sync_wrapped(*args, **kwargs):
|
||||||
|
close_old_connections()
|
||||||
|
try:
|
||||||
|
return func(*args, **kwargs)
|
||||||
|
finally:
|
||||||
|
close_old_connections()
|
||||||
|
|
||||||
|
wrapped = sync_wrapped
|
||||||
|
|
||||||
|
setattr(wrapped, _CONNECTION_CLEANUP_WRAPPED, True)
|
||||||
|
task.func = wrapped
|
||||||
|
|
||||||
|
|
||||||
def on_app_ready(app: procrastinate.App):
|
def on_app_ready(app: procrastinate.App):
|
||||||
"""This function is ran upon procrastinate initialization."""
|
"""This function is ran upon procrastinate initialization."""
|
||||||
...
|
for task in set(app.tasks.values()):
|
||||||
|
_wrap_task_with_django_connection_cleanup(task)
|
||||||
|
|||||||
@@ -0,0 +1 @@
|
|||||||
|
|
||||||
@@ -0,0 +1,298 @@
|
|||||||
|
import os
|
||||||
|
import json
|
||||||
|
from io import StringIO
|
||||||
|
from unittest.mock import patch
|
||||||
|
|
||||||
|
from django.contrib.auth import get_user_model
|
||||||
|
from django.contrib.auth.hashers import check_password
|
||||||
|
from django.core.management import call_command
|
||||||
|
from django.test import SimpleTestCase, TestCase, override_settings
|
||||||
|
from django.utils import timezone
|
||||||
|
from django.urls import reverse
|
||||||
|
from oauth2_provider.models import get_application_model
|
||||||
|
|
||||||
|
from apps.users.models import APIToken
|
||||||
|
|
||||||
|
Application = get_application_model()
|
||||||
|
|
||||||
|
|
||||||
|
@override_settings(
|
||||||
|
PUBLIC_BASE_URL="https://wygiwyh.example.com",
|
||||||
|
SECRET_KEY="test-secret-key",
|
||||||
|
OAUTH2_PROVIDER={"SCOPES": {"mcp": "Access WYGIWYH from MCP clients."}},
|
||||||
|
)
|
||||||
|
class AuthorizationServerMetadataTests(SimpleTestCase):
|
||||||
|
@override_settings(OAUTH2_DCR_ENABLED=True)
|
||||||
|
def test_returns_oauth_authorization_server_metadata(self):
|
||||||
|
response = self.client.get(reverse("oauth-authorization-server-metadata"))
|
||||||
|
|
||||||
|
self.assertEqual(response.status_code, 200)
|
||||||
|
self.assertEqual(response.json()["issuer"], "https://wygiwyh.example.com")
|
||||||
|
self.assertEqual(
|
||||||
|
response.json()["authorization_endpoint"],
|
||||||
|
"https://wygiwyh.example.com/oauth/authorize/",
|
||||||
|
)
|
||||||
|
self.assertEqual(
|
||||||
|
response.json()["registration_endpoint"],
|
||||||
|
"https://wygiwyh.example.com/oauth/register/",
|
||||||
|
)
|
||||||
|
self.assertEqual(response.json()["scopes_supported"], ["mcp"])
|
||||||
|
self.assertIn("none", response.json()["token_endpoint_auth_methods_supported"])
|
||||||
|
|
||||||
|
@override_settings(OAUTH2_DCR_ENABLED=False)
|
||||||
|
def test_omits_registration_endpoint_when_dcr_disabled(self):
|
||||||
|
response = self.client.get(reverse("oauth-authorization-server-metadata"))
|
||||||
|
|
||||||
|
self.assertEqual(response.status_code, 200)
|
||||||
|
self.assertNotIn("registration_endpoint", response.json())
|
||||||
|
|
||||||
|
|
||||||
|
@override_settings(
|
||||||
|
PUBLIC_BASE_URL="https://wygiwyh.example.com",
|
||||||
|
SECRET_KEY="test-secret-key",
|
||||||
|
OAUTH2_PROVIDER={"SCOPES": {"mcp": "Access WYGIWYH from MCP clients."}},
|
||||||
|
OAUTH2_DCR_ENABLED=True,
|
||||||
|
OAUTH2_DCR_INITIAL_ACCESS_TOKEN="",
|
||||||
|
)
|
||||||
|
class DynamicClientRegistrationTests(TestCase):
|
||||||
|
def test_registers_public_client_for_pkce_flow(self):
|
||||||
|
response = self.client.post(
|
||||||
|
reverse("oauth-dynamic-client-registration"),
|
||||||
|
data=json.dumps(
|
||||||
|
{
|
||||||
|
"client_name": "Copilot MCP",
|
||||||
|
"redirect_uris": ["http://127.0.0.1:8765/callback"],
|
||||||
|
"grant_types": ["authorization_code", "refresh_token"],
|
||||||
|
"response_types": ["code"],
|
||||||
|
"scope": "mcp",
|
||||||
|
"token_endpoint_auth_method": "none",
|
||||||
|
}
|
||||||
|
),
|
||||||
|
content_type="application/json",
|
||||||
|
)
|
||||||
|
|
||||||
|
self.assertEqual(response.status_code, 201)
|
||||||
|
payload = response.json()
|
||||||
|
self.assertEqual(payload["client_name"], "Copilot MCP")
|
||||||
|
self.assertEqual(
|
||||||
|
payload["redirect_uris"],
|
||||||
|
["http://127.0.0.1:8765/callback"],
|
||||||
|
)
|
||||||
|
self.assertEqual(
|
||||||
|
payload["grant_types"],
|
||||||
|
["authorization_code", "refresh_token"],
|
||||||
|
)
|
||||||
|
self.assertEqual(payload["response_types"], ["code"])
|
||||||
|
self.assertEqual(payload["scope"], "mcp")
|
||||||
|
self.assertEqual(payload["token_endpoint_auth_method"], "none")
|
||||||
|
self.assertNotIn("client_secret", payload)
|
||||||
|
|
||||||
|
application = Application.objects.get(client_id=payload["client_id"])
|
||||||
|
self.assertEqual(application.name, "Copilot MCP")
|
||||||
|
self.assertEqual(application.client_type, Application.CLIENT_PUBLIC)
|
||||||
|
self.assertEqual(
|
||||||
|
application.authorization_grant_type,
|
||||||
|
Application.GRANT_AUTHORIZATION_CODE,
|
||||||
|
)
|
||||||
|
self.assertEqual(
|
||||||
|
application.redirect_uris,
|
||||||
|
"http://127.0.0.1:8765/callback",
|
||||||
|
)
|
||||||
|
|
||||||
|
def test_registers_confidential_client_with_generated_secret(self):
|
||||||
|
response = self.client.post(
|
||||||
|
reverse("oauth-dynamic-client-registration"),
|
||||||
|
data=json.dumps(
|
||||||
|
{
|
||||||
|
"client_name": "Confidential MCP",
|
||||||
|
"redirect_uris": ["http://127.0.0.1:8765/callback"],
|
||||||
|
"token_endpoint_auth_method": "client_secret_basic",
|
||||||
|
}
|
||||||
|
),
|
||||||
|
content_type="application/json",
|
||||||
|
)
|
||||||
|
|
||||||
|
self.assertEqual(response.status_code, 201)
|
||||||
|
payload = response.json()
|
||||||
|
self.assertEqual(payload["token_endpoint_auth_method"], "client_secret_basic")
|
||||||
|
self.assertEqual(payload["scope"], "mcp")
|
||||||
|
self.assertEqual(payload["client_secret_expires_at"], 0)
|
||||||
|
self.assertTrue(payload["client_secret"])
|
||||||
|
|
||||||
|
application = Application.objects.get(client_id=payload["client_id"])
|
||||||
|
self.assertEqual(application.client_type, Application.CLIENT_CONFIDENTIAL)
|
||||||
|
self.assertTrue(check_password(payload["client_secret"], application.client_secret))
|
||||||
|
|
||||||
|
def test_rejects_unsupported_token_auth_method(self):
|
||||||
|
response = self.client.post(
|
||||||
|
reverse("oauth-dynamic-client-registration"),
|
||||||
|
data=json.dumps(
|
||||||
|
{
|
||||||
|
"redirect_uris": ["http://127.0.0.1:8765/callback"],
|
||||||
|
"token_endpoint_auth_method": "private_key_jwt",
|
||||||
|
}
|
||||||
|
),
|
||||||
|
content_type="application/json",
|
||||||
|
)
|
||||||
|
|
||||||
|
self.assertEqual(response.status_code, 400)
|
||||||
|
self.assertEqual(response.json()["error"], "invalid_client_metadata")
|
||||||
|
self.assertIn("token_endpoint_auth_method", response.json()["error_description"])
|
||||||
|
|
||||||
|
def test_rejects_missing_redirect_uris(self):
|
||||||
|
response = self.client.post(
|
||||||
|
reverse("oauth-dynamic-client-registration"),
|
||||||
|
data=json.dumps({"client_name": "No redirect"}),
|
||||||
|
content_type="application/json",
|
||||||
|
)
|
||||||
|
|
||||||
|
self.assertEqual(response.status_code, 400)
|
||||||
|
self.assertEqual(response.json()["error"], "invalid_client_metadata")
|
||||||
|
self.assertIn("redirect_uris", response.json()["error_description"])
|
||||||
|
|
||||||
|
@override_settings(OAUTH2_DCR_ENABLED=False)
|
||||||
|
def test_returns_404_when_dcr_disabled(self):
|
||||||
|
response = self.client.post(
|
||||||
|
reverse("oauth-dynamic-client-registration"),
|
||||||
|
data=json.dumps({"redirect_uris": ["http://127.0.0.1:8765/callback"]}),
|
||||||
|
content_type="application/json",
|
||||||
|
)
|
||||||
|
|
||||||
|
self.assertEqual(response.status_code, 404)
|
||||||
|
self.assertEqual(Application.objects.count(), 0)
|
||||||
|
|
||||||
|
|
||||||
|
@override_settings(
|
||||||
|
PUBLIC_BASE_URL="https://wygiwyh.example.com",
|
||||||
|
SECRET_KEY="test-secret-key",
|
||||||
|
OAUTH2_PROVIDER={"SCOPES": {"mcp": "Access WYGIWYH from MCP clients."}},
|
||||||
|
OAUTH2_DCR_ENABLED=True,
|
||||||
|
OAUTH2_DCR_INITIAL_ACCESS_TOKEN="s3cret-iat",
|
||||||
|
)
|
||||||
|
class DynamicClientRegistrationInitialAccessTokenTests(TestCase):
|
||||||
|
def test_rejects_registration_without_initial_access_token(self):
|
||||||
|
response = self.client.post(
|
||||||
|
reverse("oauth-dynamic-client-registration"),
|
||||||
|
data=json.dumps({"redirect_uris": ["http://127.0.0.1:8765/callback"]}),
|
||||||
|
content_type="application/json",
|
||||||
|
)
|
||||||
|
|
||||||
|
self.assertEqual(response.status_code, 401)
|
||||||
|
self.assertEqual(response.json()["error"], "invalid_token")
|
||||||
|
self.assertEqual(Application.objects.count(), 0)
|
||||||
|
|
||||||
|
def test_allows_registration_with_initial_access_token(self):
|
||||||
|
response = self.client.post(
|
||||||
|
reverse("oauth-dynamic-client-registration"),
|
||||||
|
data=json.dumps(
|
||||||
|
{
|
||||||
|
"redirect_uris": ["http://127.0.0.1:8765/callback"],
|
||||||
|
"token_endpoint_auth_method": "none",
|
||||||
|
}
|
||||||
|
),
|
||||||
|
content_type="application/json",
|
||||||
|
HTTP_AUTHORIZATION="Bearer s3cret-iat",
|
||||||
|
)
|
||||||
|
|
||||||
|
self.assertEqual(response.status_code, 201)
|
||||||
|
self.assertEqual(Application.objects.count(), 1)
|
||||||
|
|
||||||
|
|
||||||
|
class SetupOAuthCommandTests(TestCase):
|
||||||
|
@patch.dict(
|
||||||
|
os.environ,
|
||||||
|
{
|
||||||
|
"MCP_OAUTH_CLIENT_ID": "mcp-wygiwyh",
|
||||||
|
"MCP_OAUTH_CLIENT_SECRET": "super-secret",
|
||||||
|
"MCP_OAUTH_REDIRECT_URIS": "http://127.0.0.1:8765/callback",
|
||||||
|
},
|
||||||
|
clear=False,
|
||||||
|
)
|
||||||
|
def test_creates_mcp_oauth_application(self):
|
||||||
|
call_command("setup_oauth")
|
||||||
|
|
||||||
|
application = Application.objects.get(client_id="mcp-wygiwyh")
|
||||||
|
self.assertEqual(application.name, "WYGIWYH MCP")
|
||||||
|
self.assertEqual(application.client_type, Application.CLIENT_CONFIDENTIAL)
|
||||||
|
self.assertEqual(
|
||||||
|
application.authorization_grant_type,
|
||||||
|
Application.GRANT_AUTHORIZATION_CODE,
|
||||||
|
)
|
||||||
|
self.assertEqual(
|
||||||
|
application.redirect_uris,
|
||||||
|
"http://127.0.0.1:8765/callback",
|
||||||
|
)
|
||||||
|
self.assertFalse(application.skip_authorization)
|
||||||
|
self.assertTrue(check_password("super-secret", application.client_secret))
|
||||||
|
|
||||||
|
@patch.dict(
|
||||||
|
os.environ,
|
||||||
|
{
|
||||||
|
"MCP_OAUTH_CLIENT_ID": "mcp-wygiwyh",
|
||||||
|
"MCP_OAUTH_CLIENT_SECRET": "new-secret",
|
||||||
|
"MCP_OAUTH_REDIRECT_URIS": "http://127.0.0.1:8765/callback http://localhost:8765/callback",
|
||||||
|
"MCP_OAUTH_CLIENT_NAME": "WYGIWYH MCP Production",
|
||||||
|
"MCP_OAUTH_SKIP_AUTHORIZATION": "true",
|
||||||
|
},
|
||||||
|
clear=False,
|
||||||
|
)
|
||||||
|
def test_updates_existing_mcp_oauth_application(self):
|
||||||
|
Application.objects.create(
|
||||||
|
client_id="mcp-wygiwyh",
|
||||||
|
client_secret="old-secret",
|
||||||
|
name="Old Name",
|
||||||
|
client_type=Application.CLIENT_CONFIDENTIAL,
|
||||||
|
authorization_grant_type=Application.GRANT_AUTHORIZATION_CODE,
|
||||||
|
redirect_uris="http://127.0.0.1:8765/callback",
|
||||||
|
skip_authorization=False,
|
||||||
|
)
|
||||||
|
|
||||||
|
call_command("setup_oauth")
|
||||||
|
|
||||||
|
application = Application.objects.get(client_id="mcp-wygiwyh")
|
||||||
|
self.assertEqual(application.name, "WYGIWYH MCP Production")
|
||||||
|
self.assertEqual(
|
||||||
|
application.redirect_uris,
|
||||||
|
"http://127.0.0.1:8765/callback http://localhost:8765/callback",
|
||||||
|
)
|
||||||
|
self.assertTrue(application.skip_authorization)
|
||||||
|
self.assertTrue(check_password("new-secret", application.client_secret))
|
||||||
|
|
||||||
|
|
||||||
|
class CreateAPITokenCommandTests(TestCase):
|
||||||
|
def setUp(self):
|
||||||
|
self.user = get_user_model().objects.create_user(
|
||||||
|
email="n8n@example.com",
|
||||||
|
password="test-password",
|
||||||
|
)
|
||||||
|
|
||||||
|
def test_creates_hashed_api_token_and_prints_raw_value(self):
|
||||||
|
stdout = StringIO()
|
||||||
|
|
||||||
|
call_command(
|
||||||
|
"create_api_token",
|
||||||
|
self.user.email,
|
||||||
|
"--name",
|
||||||
|
"n8n sync",
|
||||||
|
stdout=stdout,
|
||||||
|
)
|
||||||
|
|
||||||
|
token = APIToken.objects.get(user=self.user, name="n8n sync")
|
||||||
|
lines = [line.strip() for line in stdout.getvalue().splitlines() if line.strip()]
|
||||||
|
raw_token = lines[-1]
|
||||||
|
|
||||||
|
self.assertTrue(raw_token.startswith(APIToken.TOKEN_PREFIX))
|
||||||
|
self.assertNotEqual(token.token_hash, raw_token)
|
||||||
|
self.assertTrue(token.check_secret(APIToken.parse_raw_token(raw_token)[1]))
|
||||||
|
|
||||||
|
def test_supports_expiring_tokens(self):
|
||||||
|
call_command(
|
||||||
|
"create_api_token",
|
||||||
|
self.user.email,
|
||||||
|
"--expires-in-days",
|
||||||
|
"7",
|
||||||
|
)
|
||||||
|
|
||||||
|
token = APIToken.objects.get(user=self.user)
|
||||||
|
self.assertIsNotNone(token.expires_at)
|
||||||
|
self.assertGreater(token.expires_at, timezone.now())
|
||||||
@@ -0,0 +1,269 @@
|
|||||||
|
from django.contrib.auth import get_user_model
|
||||||
|
from django.contrib.auth.models import AnonymousUser
|
||||||
|
from django.core.exceptions import PermissionDenied
|
||||||
|
from django.http import Http404
|
||||||
|
from django.test import RequestFactory, TestCase, override_settings
|
||||||
|
|
||||||
|
from apps.common.functions.permissions import (
|
||||||
|
EDIT,
|
||||||
|
READ,
|
||||||
|
get_shared_object_or_error,
|
||||||
|
)
|
||||||
|
from apps.common.middleware.thread_local import delete_current_user, write_current_user
|
||||||
|
from apps.rules.models import TransactionRule, TransactionRuleAction
|
||||||
|
|
||||||
|
|
||||||
|
@override_settings(
|
||||||
|
STORAGES={
|
||||||
|
"default": {"BACKEND": "django.core.files.storage.FileSystemStorage"},
|
||||||
|
"staticfiles": {
|
||||||
|
"BACKEND": "django.contrib.staticfiles.storage.StaticFilesStorage"
|
||||||
|
},
|
||||||
|
},
|
||||||
|
WHITENOISE_AUTOREFRESH=True,
|
||||||
|
)
|
||||||
|
class SharedObjectPredicateTests(TestCase):
|
||||||
|
"""Unit tests for is_visible_to / is_editable_by on SharedObject."""
|
||||||
|
|
||||||
|
def setUp(self):
|
||||||
|
User = get_user_model()
|
||||||
|
self.owner = User.objects.create_user(
|
||||||
|
email="owner@test.com", password="testpass123"
|
||||||
|
)
|
||||||
|
self.shared_user = User.objects.create_user(
|
||||||
|
email="shared@test.com", password="testpass123"
|
||||||
|
)
|
||||||
|
self.stranger = User.objects.create_user(
|
||||||
|
email="stranger@test.com", password="testpass123"
|
||||||
|
)
|
||||||
|
|
||||||
|
def _rule(self, **kwargs):
|
||||||
|
kwargs.setdefault("name", "Rule")
|
||||||
|
kwargs.setdefault("trigger", "True")
|
||||||
|
return TransactionRule.all_objects.create(**kwargs)
|
||||||
|
|
||||||
|
def test_owner_can_read_and_edit_own_private_rule(self):
|
||||||
|
rule = self._rule(owner=self.owner, visibility="private")
|
||||||
|
|
||||||
|
self.assertTrue(rule.is_visible_to(self.owner))
|
||||||
|
self.assertTrue(rule.is_editable_by(self.owner))
|
||||||
|
|
||||||
|
def test_private_rule_is_invisible_to_stranger(self):
|
||||||
|
rule = self._rule(owner=self.owner, visibility="private")
|
||||||
|
|
||||||
|
self.assertFalse(rule.is_visible_to(self.stranger))
|
||||||
|
self.assertFalse(rule.is_editable_by(self.stranger))
|
||||||
|
|
||||||
|
def test_shared_rule_is_readable_but_not_editable(self):
|
||||||
|
rule = self._rule(owner=self.owner, visibility="private")
|
||||||
|
rule.shared_with.add(self.shared_user)
|
||||||
|
|
||||||
|
self.assertTrue(rule.is_visible_to(self.shared_user))
|
||||||
|
self.assertFalse(rule.is_editable_by(self.shared_user))
|
||||||
|
|
||||||
|
def test_public_rule_is_readable_but_not_editable(self):
|
||||||
|
rule = self._rule(owner=self.owner, visibility="public")
|
||||||
|
|
||||||
|
self.assertTrue(rule.is_visible_to(self.stranger))
|
||||||
|
self.assertFalse(rule.is_editable_by(self.stranger))
|
||||||
|
|
||||||
|
def test_unowned_rule_stays_readable_and_editable_by_everyone(self):
|
||||||
|
rule = self._rule(owner=None, visibility="private")
|
||||||
|
|
||||||
|
self.assertTrue(rule.is_visible_to(self.stranger))
|
||||||
|
self.assertTrue(rule.is_editable_by(self.stranger))
|
||||||
|
|
||||||
|
def test_anonymous_user_gets_no_access_to_owned_rules(self):
|
||||||
|
rule = self._rule(owner=self.owner, visibility="private")
|
||||||
|
|
||||||
|
self.assertFalse(rule.is_visible_to(AnonymousUser()))
|
||||||
|
self.assertFalse(rule.is_editable_by(AnonymousUser()))
|
||||||
|
|
||||||
|
|
||||||
|
@override_settings(
|
||||||
|
STORAGES={
|
||||||
|
"default": {"BACKEND": "django.core.files.storage.FileSystemStorage"},
|
||||||
|
"staticfiles": {
|
||||||
|
"BACKEND": "django.contrib.staticfiles.storage.StaticFilesStorage"
|
||||||
|
},
|
||||||
|
},
|
||||||
|
WHITENOISE_AUTOREFRESH=True,
|
||||||
|
)
|
||||||
|
class SharedObjectPredicateParityTests(TestCase):
|
||||||
|
"""is_visible_to must agree with what SharedObjectManager returns.
|
||||||
|
|
||||||
|
The manager filters at the queryset level and the predicate checks a single
|
||||||
|
instance, so the two cannot share an implementation. This asserts they do
|
||||||
|
not drift apart.
|
||||||
|
"""
|
||||||
|
|
||||||
|
def setUp(self):
|
||||||
|
User = get_user_model()
|
||||||
|
self.owner = User.objects.create_user(
|
||||||
|
email="owner@test.com", password="testpass123"
|
||||||
|
)
|
||||||
|
self.other = User.objects.create_user(
|
||||||
|
email="other@test.com", password="testpass123"
|
||||||
|
)
|
||||||
|
self.addCleanup(self._clear_current_user)
|
||||||
|
|
||||||
|
def _clear_current_user(self):
|
||||||
|
try:
|
||||||
|
delete_current_user()
|
||||||
|
except AttributeError:
|
||||||
|
pass
|
||||||
|
|
||||||
|
def test_manager_and_predicate_agree_over_every_combination(self):
|
||||||
|
combinations = []
|
||||||
|
for owner in (self.owner, self.other, None):
|
||||||
|
for visibility in ("private", "public"):
|
||||||
|
for shared in (True, False):
|
||||||
|
rule = TransactionRule.all_objects.create(
|
||||||
|
name=f"{owner}-{visibility}-{shared}",
|
||||||
|
trigger="True",
|
||||||
|
owner=owner,
|
||||||
|
visibility=visibility,
|
||||||
|
)
|
||||||
|
if shared:
|
||||||
|
rule.shared_with.add(self.owner)
|
||||||
|
combinations.append(rule)
|
||||||
|
|
||||||
|
write_current_user(self.owner)
|
||||||
|
visible_ids = set(TransactionRule.objects.values_list("id", flat=True))
|
||||||
|
|
||||||
|
for rule in combinations:
|
||||||
|
with self.subTest(rule=rule.name):
|
||||||
|
self.assertEqual(
|
||||||
|
rule.id in visible_ids,
|
||||||
|
rule.is_visible_to(self.owner),
|
||||||
|
f"manager and is_visible_to disagree for {rule.name}",
|
||||||
|
)
|
||||||
|
|
||||||
|
|
||||||
|
@override_settings(
|
||||||
|
STORAGES={
|
||||||
|
"default": {"BACKEND": "django.core.files.storage.FileSystemStorage"},
|
||||||
|
"staticfiles": {
|
||||||
|
"BACKEND": "django.contrib.staticfiles.storage.StaticFilesStorage"
|
||||||
|
},
|
||||||
|
},
|
||||||
|
WHITENOISE_AUTOREFRESH=True,
|
||||||
|
)
|
||||||
|
class GetSharedObjectOrErrorTests(TestCase):
|
||||||
|
def setUp(self):
|
||||||
|
User = get_user_model()
|
||||||
|
self.owner = User.objects.create_user(
|
||||||
|
email="owner@test.com", password="testpass123"
|
||||||
|
)
|
||||||
|
self.stranger = User.objects.create_user(
|
||||||
|
email="stranger@test.com", password="testpass123"
|
||||||
|
)
|
||||||
|
self.factory = RequestFactory()
|
||||||
|
|
||||||
|
self.public_rule = TransactionRule.all_objects.create(
|
||||||
|
name="Public", trigger="True", owner=self.owner, visibility="public"
|
||||||
|
)
|
||||||
|
self.private_rule = TransactionRule.all_objects.create(
|
||||||
|
name="Private", trigger="True", owner=self.owner, visibility="private"
|
||||||
|
)
|
||||||
|
self.public_action = TransactionRuleAction.objects.create(
|
||||||
|
rule=self.public_rule, field="notes", value="x"
|
||||||
|
)
|
||||||
|
self.private_action = TransactionRuleAction.objects.create(
|
||||||
|
rule=self.private_rule, field="notes", value="x"
|
||||||
|
)
|
||||||
|
|
||||||
|
self.addCleanup(self._clear_current_user)
|
||||||
|
|
||||||
|
def _clear_current_user(self):
|
||||||
|
try:
|
||||||
|
delete_current_user()
|
||||||
|
except AttributeError:
|
||||||
|
pass
|
||||||
|
|
||||||
|
def _request(self, user):
|
||||||
|
request = self.factory.get("/")
|
||||||
|
request.user = user
|
||||||
|
write_current_user(user)
|
||||||
|
return request
|
||||||
|
|
||||||
|
def test_read_allows_visible_object(self):
|
||||||
|
request = self._request(self.stranger)
|
||||||
|
|
||||||
|
rule = get_shared_object_or_error(
|
||||||
|
TransactionRule, request, id=self.public_rule.id, level=READ
|
||||||
|
)
|
||||||
|
|
||||||
|
self.assertEqual(rule, self.public_rule)
|
||||||
|
|
||||||
|
def test_edit_denies_visible_but_unowned_object_with_403(self):
|
||||||
|
request = self._request(self.stranger)
|
||||||
|
|
||||||
|
with self.assertRaises(PermissionDenied):
|
||||||
|
get_shared_object_or_error(
|
||||||
|
TransactionRule, request, id=self.public_rule.id, level=EDIT
|
||||||
|
)
|
||||||
|
|
||||||
|
def test_edit_allows_owner(self):
|
||||||
|
request = self._request(self.owner)
|
||||||
|
|
||||||
|
rule = get_shared_object_or_error(
|
||||||
|
TransactionRule, request, id=self.public_rule.id, level=EDIT
|
||||||
|
)
|
||||||
|
|
||||||
|
self.assertEqual(rule, self.public_rule)
|
||||||
|
|
||||||
|
def test_invisible_object_raises_404_not_403(self):
|
||||||
|
"""403 must not confirm the existence of an object the user cannot see."""
|
||||||
|
request = self._request(self.stranger)
|
||||||
|
|
||||||
|
with self.assertRaises(Http404):
|
||||||
|
get_shared_object_or_error(
|
||||||
|
TransactionRule, request, id=self.private_rule.id, level=EDIT
|
||||||
|
)
|
||||||
|
|
||||||
|
def test_via_traverses_to_the_governing_object(self):
|
||||||
|
request = self._request(self.stranger)
|
||||||
|
|
||||||
|
with self.assertRaises(PermissionDenied):
|
||||||
|
get_shared_object_or_error(
|
||||||
|
TransactionRuleAction,
|
||||||
|
request,
|
||||||
|
id=self.public_action.id,
|
||||||
|
level=EDIT,
|
||||||
|
via="rule",
|
||||||
|
)
|
||||||
|
|
||||||
|
def test_via_hides_children_of_invisible_parents_behind_404(self):
|
||||||
|
"""TransactionRuleAction has an unscoped manager, so the id is reachable."""
|
||||||
|
request = self._request(self.stranger)
|
||||||
|
|
||||||
|
with self.assertRaises(Http404):
|
||||||
|
get_shared_object_or_error(
|
||||||
|
TransactionRuleAction,
|
||||||
|
request,
|
||||||
|
id=self.private_action.id,
|
||||||
|
level=EDIT,
|
||||||
|
via="rule",
|
||||||
|
)
|
||||||
|
|
||||||
|
def test_unresolvable_via_path_fails_closed(self):
|
||||||
|
"""A typo'd path must raise, never silently grant access."""
|
||||||
|
request = self._request(self.stranger)
|
||||||
|
|
||||||
|
with self.assertRaises(AttributeError):
|
||||||
|
get_shared_object_or_error(
|
||||||
|
TransactionRuleAction,
|
||||||
|
request,
|
||||||
|
id=self.public_action.id,
|
||||||
|
level=EDIT,
|
||||||
|
via="rulee",
|
||||||
|
)
|
||||||
|
|
||||||
|
def test_unknown_level_is_rejected(self):
|
||||||
|
request = self._request(self.owner)
|
||||||
|
|
||||||
|
with self.assertRaises(ValueError):
|
||||||
|
get_shared_object_or_error(
|
||||||
|
TransactionRule, request, id=self.public_rule.id, level="write"
|
||||||
|
)
|
||||||
@@ -0,0 +1,89 @@
|
|||||||
|
from unittest.mock import patch
|
||||||
|
|
||||||
|
import procrastinate
|
||||||
|
from django.db import connection
|
||||||
|
from django.test import SimpleTestCase, TransactionTestCase
|
||||||
|
from procrastinate.testing import InMemoryConnector
|
||||||
|
|
||||||
|
from apps.common.procrastinate import on_app_ready
|
||||||
|
|
||||||
|
|
||||||
|
def make_app_with_task(func):
|
||||||
|
app = procrastinate.App(connector=InMemoryConnector())
|
||||||
|
task = app.task(name="sample_task")(func)
|
||||||
|
|
||||||
|
return app, task
|
||||||
|
|
||||||
|
|
||||||
|
class ProcrastinateConnectionCleanupTests(SimpleTestCase):
|
||||||
|
def test_app_ready_closes_old_connections_around_sync_tasks(self):
|
||||||
|
calls = []
|
||||||
|
|
||||||
|
def sample_task(value):
|
||||||
|
calls.append(("task", value))
|
||||||
|
return value * 2
|
||||||
|
|
||||||
|
app, task = make_app_with_task(sample_task)
|
||||||
|
|
||||||
|
with patch(
|
||||||
|
"apps.common.procrastinate.close_old_connections",
|
||||||
|
create=True,
|
||||||
|
side_effect=lambda: calls.append(("cleanup", None)),
|
||||||
|
):
|
||||||
|
on_app_ready(app)
|
||||||
|
|
||||||
|
result = task.func(3)
|
||||||
|
|
||||||
|
self.assertEqual(result, 6)
|
||||||
|
self.assertEqual(
|
||||||
|
calls,
|
||||||
|
[
|
||||||
|
("cleanup", None),
|
||||||
|
("task", 3),
|
||||||
|
("cleanup", None),
|
||||||
|
],
|
||||||
|
)
|
||||||
|
|
||||||
|
def test_app_ready_closes_old_connections_when_sync_task_raises(self):
|
||||||
|
calls = []
|
||||||
|
|
||||||
|
def sample_task():
|
||||||
|
calls.append(("task", None))
|
||||||
|
raise RuntimeError("boom")
|
||||||
|
|
||||||
|
app, task = make_app_with_task(sample_task)
|
||||||
|
|
||||||
|
with patch(
|
||||||
|
"apps.common.procrastinate.close_old_connections",
|
||||||
|
create=True,
|
||||||
|
side_effect=lambda: calls.append(("cleanup", None)),
|
||||||
|
):
|
||||||
|
on_app_ready(app)
|
||||||
|
|
||||||
|
with self.assertRaises(RuntimeError):
|
||||||
|
task.func()
|
||||||
|
|
||||||
|
self.assertEqual(
|
||||||
|
calls,
|
||||||
|
[
|
||||||
|
("cleanup", None),
|
||||||
|
("task", None),
|
||||||
|
("cleanup", None),
|
||||||
|
],
|
||||||
|
)
|
||||||
|
|
||||||
|
|
||||||
|
class ProcrastinateConnectionRecoveryTests(TransactionTestCase):
|
||||||
|
def test_wrapped_task_recovers_from_closed_django_connection(self):
|
||||||
|
def sample_task():
|
||||||
|
with connection.cursor() as cursor:
|
||||||
|
cursor.execute("SELECT 1")
|
||||||
|
return cursor.fetchone()[0]
|
||||||
|
|
||||||
|
app, task = make_app_with_task(sample_task)
|
||||||
|
on_app_ready(app)
|
||||||
|
|
||||||
|
connection.ensure_connection()
|
||||||
|
connection.connection.close()
|
||||||
|
|
||||||
|
self.assertEqual(task.func(), 1)
|
||||||
@@ -0,0 +1,161 @@
|
|||||||
|
"""Delete semantics for every SharedObject-backed model.
|
||||||
|
|
||||||
|
Each of these views carried the same inverted condition::
|
||||||
|
|
||||||
|
if obj.owner != request.user and request.user in obj.shared_with.all():
|
||||||
|
obj.shared_with.remove(request.user) # unshare
|
||||||
|
else:
|
||||||
|
obj.delete() # <- public objects landed here
|
||||||
|
|
||||||
|
so an object its owner had made public could be destroyed by any authenticated
|
||||||
|
user. The rule is: the owner deletes, a shared user revokes only their own
|
||||||
|
access, and nobody else may do either.
|
||||||
|
"""
|
||||||
|
|
||||||
|
from django.contrib.auth import get_user_model
|
||||||
|
from django.test import TestCase, override_settings
|
||||||
|
from django.urls import reverse
|
||||||
|
|
||||||
|
from apps.accounts.models import Account, AccountGroup
|
||||||
|
from apps.currencies.models import Currency
|
||||||
|
from apps.dca.models import DCAStrategy
|
||||||
|
from apps.rules.models import TransactionRule
|
||||||
|
from apps.transactions.models import (
|
||||||
|
TransactionCategory,
|
||||||
|
TransactionEntity,
|
||||||
|
TransactionTag,
|
||||||
|
)
|
||||||
|
|
||||||
|
HTMX = {"HTTP_HX_REQUEST": "true"}
|
||||||
|
|
||||||
|
|
||||||
|
@override_settings(
|
||||||
|
STORAGES={
|
||||||
|
"default": {"BACKEND": "django.core.files.storage.FileSystemStorage"},
|
||||||
|
"staticfiles": {
|
||||||
|
"BACKEND": "django.contrib.staticfiles.storage.StaticFilesStorage"
|
||||||
|
},
|
||||||
|
},
|
||||||
|
WHITENOISE_AUTOREFRESH=True,
|
||||||
|
DEMO=False,
|
||||||
|
)
|
||||||
|
class SharedObjectDeletionTests(TestCase):
|
||||||
|
def setUp(self):
|
||||||
|
User = get_user_model()
|
||||||
|
self.owner = User.objects.create_user(
|
||||||
|
email="owner@test.com", password="testpass123"
|
||||||
|
)
|
||||||
|
self.shared_user = User.objects.create_user(
|
||||||
|
email="shared@test.com", password="testpass123"
|
||||||
|
)
|
||||||
|
self.stranger = User.objects.create_user(
|
||||||
|
email="stranger@test.com", password="testpass123"
|
||||||
|
)
|
||||||
|
self.currency = Currency.objects.create(
|
||||||
|
code="USD", name="US Dollar", decimal_places=2
|
||||||
|
)
|
||||||
|
|
||||||
|
def cases(self):
|
||||||
|
"""(label, model, delete url name, url kwarg, extra create kwargs)."""
|
||||||
|
return [
|
||||||
|
(
|
||||||
|
"account",
|
||||||
|
Account,
|
||||||
|
"account_delete",
|
||||||
|
"pk",
|
||||||
|
{"currency": self.currency},
|
||||||
|
),
|
||||||
|
("account group", AccountGroup, "account_group_delete", "pk", {}),
|
||||||
|
(
|
||||||
|
"category",
|
||||||
|
TransactionCategory,
|
||||||
|
"category_delete",
|
||||||
|
"category_id",
|
||||||
|
{},
|
||||||
|
),
|
||||||
|
("tag", TransactionTag, "tag_delete", "tag_id", {}),
|
||||||
|
("entity", TransactionEntity, "entity_delete", "entity_id", {}),
|
||||||
|
(
|
||||||
|
"rule",
|
||||||
|
TransactionRule,
|
||||||
|
"transaction_rule_delete",
|
||||||
|
"transaction_rule_id",
|
||||||
|
{"trigger": "True"},
|
||||||
|
),
|
||||||
|
(
|
||||||
|
"dca strategy",
|
||||||
|
DCAStrategy,
|
||||||
|
"dca_strategy_delete",
|
||||||
|
"strategy_id",
|
||||||
|
{
|
||||||
|
"target_currency": self.currency,
|
||||||
|
"payment_currency": self.currency,
|
||||||
|
},
|
||||||
|
),
|
||||||
|
]
|
||||||
|
|
||||||
|
def make(self, model, visibility, extra):
|
||||||
|
return model.all_objects.create(
|
||||||
|
name="Target", owner=self.owner, visibility=visibility, **extra
|
||||||
|
)
|
||||||
|
|
||||||
|
def delete(self, url_name, kwarg, obj):
|
||||||
|
return self.client.delete(reverse(url_name, kwargs={kwarg: obj.id}), **HTMX)
|
||||||
|
|
||||||
|
def test_stranger_cannot_delete_a_public_object(self):
|
||||||
|
for label, model, url_name, kwarg, extra in self.cases():
|
||||||
|
with self.subTest(model=label):
|
||||||
|
obj = self.make(model, "public", extra)
|
||||||
|
self.client.force_login(self.stranger)
|
||||||
|
|
||||||
|
response = self.delete(url_name, kwarg, obj)
|
||||||
|
|
||||||
|
self.assertEqual(response.status_code, 403, label)
|
||||||
|
self.assertTrue(
|
||||||
|
model.all_objects.filter(pk=obj.pk).exists(),
|
||||||
|
f"{label} was deleted by a non-owner",
|
||||||
|
)
|
||||||
|
|
||||||
|
def test_shared_user_deleting_only_revokes_their_own_access(self):
|
||||||
|
for label, model, url_name, kwarg, extra in self.cases():
|
||||||
|
with self.subTest(model=label):
|
||||||
|
obj = self.make(model, "private", extra)
|
||||||
|
obj.shared_with.add(self.shared_user)
|
||||||
|
self.client.force_login(self.shared_user)
|
||||||
|
|
||||||
|
response = self.delete(url_name, kwarg, obj)
|
||||||
|
|
||||||
|
self.assertEqual(response.status_code, 204, label)
|
||||||
|
self.assertTrue(
|
||||||
|
model.all_objects.filter(pk=obj.pk).exists(),
|
||||||
|
f"{label} was deleted by a shared user",
|
||||||
|
)
|
||||||
|
self.assertNotIn(self.shared_user, obj.shared_with.all())
|
||||||
|
|
||||||
|
def test_owner_can_delete(self):
|
||||||
|
for label, model, url_name, kwarg, extra in self.cases():
|
||||||
|
with self.subTest(model=label):
|
||||||
|
obj = self.make(model, "private", extra)
|
||||||
|
self.client.force_login(self.owner)
|
||||||
|
|
||||||
|
response = self.delete(url_name, kwarg, obj)
|
||||||
|
|
||||||
|
self.assertEqual(response.status_code, 204, label)
|
||||||
|
self.assertFalse(
|
||||||
|
model.all_objects.filter(pk=obj.pk).exists(),
|
||||||
|
f"{label} was not deleted by its owner",
|
||||||
|
)
|
||||||
|
|
||||||
|
def test_unowned_objects_stay_deletable(self):
|
||||||
|
"""Legacy objects with no owner are editable by everyone by design."""
|
||||||
|
for label, model, url_name, kwarg, extra in self.cases():
|
||||||
|
with self.subTest(model=label):
|
||||||
|
obj = model.all_objects.create(
|
||||||
|
name="Legacy", owner=None, visibility="private", **extra
|
||||||
|
)
|
||||||
|
self.client.force_login(self.stranger)
|
||||||
|
|
||||||
|
response = self.delete(url_name, kwarg, obj)
|
||||||
|
|
||||||
|
self.assertEqual(response.status_code, 204, label)
|
||||||
|
self.assertFalse(model.all_objects.filter(pk=obj.pk).exists(), label)
|
||||||
@@ -1,5 +1,4 @@
|
|||||||
import logging
|
import logging
|
||||||
from datetime import timedelta
|
|
||||||
|
|
||||||
from django.db.models import QuerySet
|
from django.db.models import QuerySet
|
||||||
from django.utils import timezone
|
from django.utils import timezone
|
||||||
@@ -18,6 +17,7 @@ PROVIDER_MAPPING = {
|
|||||||
"frankfurter": providers.FrankfurterProvider,
|
"frankfurter": providers.FrankfurterProvider,
|
||||||
"twelvedata": providers.TwelveDataProvider,
|
"twelvedata": providers.TwelveDataProvider,
|
||||||
"twelvedatamarkets": providers.TwelveDataMarketsProvider,
|
"twelvedatamarkets": providers.TwelveDataMarketsProvider,
|
||||||
|
"yfinance": providers.YFinanceMarketsProvider,
|
||||||
}
|
}
|
||||||
|
|
||||||
|
|
||||||
@@ -258,7 +258,10 @@ class ExchangeRateFetcher:
|
|||||||
processed_pairs.add((from_currency.id, to_currency.id))
|
processed_pairs.add((from_currency.id, to_currency.id))
|
||||||
|
|
||||||
service.last_fetch = timezone.now()
|
service.last_fetch = timezone.now()
|
||||||
|
service.failure_count = 0
|
||||||
service.save()
|
service.save()
|
||||||
|
|
||||||
except Exception as e:
|
except Exception as e:
|
||||||
logger.error(f"Error fetching rates for {service.name}: {e}")
|
logger.error(f"Error fetching rates for {service.name}: {e}")
|
||||||
|
service.failure_count += 1
|
||||||
|
service.save()
|
||||||
|
|||||||
@@ -503,3 +503,82 @@ class TwelveDataMarketsProvider(ExchangeRateProvider):
|
|||||||
)
|
)
|
||||||
|
|
||||||
return results
|
return results
|
||||||
|
|
||||||
|
|
||||||
|
class YFinanceMarketsProvider(ExchangeRateProvider):
|
||||||
|
"""Fetch market prices for Yahoo Finance symbols using yfinance.
|
||||||
|
|
||||||
|
Currency codes are passed to Yahoo Finance verbatim. For example, use
|
||||||
|
``PETR4.SA`` for Petrobras on B3 or ``AAPL`` for Apple. The configured
|
||||||
|
exchange currency is treated as the currency of the Yahoo quote.
|
||||||
|
"""
|
||||||
|
|
||||||
|
rates_inverted = True
|
||||||
|
|
||||||
|
def __init__(self, api_key: str = None, ticker_factory=None):
|
||||||
|
super().__init__(api_key)
|
||||||
|
self._ticker_factory = ticker_factory
|
||||||
|
|
||||||
|
@classmethod
|
||||||
|
def requires_api_key(cls) -> bool:
|
||||||
|
return False
|
||||||
|
|
||||||
|
def _get_ticker_factory(self):
|
||||||
|
if self._ticker_factory is None:
|
||||||
|
try:
|
||||||
|
import yfinance as yf
|
||||||
|
except ImportError as exc:
|
||||||
|
raise RuntimeError(
|
||||||
|
"The yfinance package is required for the Yahoo Finance provider."
|
||||||
|
) from exc
|
||||||
|
|
||||||
|
self._ticker_factory = yf.Ticker
|
||||||
|
|
||||||
|
return self._ticker_factory
|
||||||
|
|
||||||
|
def get_rates(
|
||||||
|
self, target_currencies: QuerySet, exchange_currencies: set
|
||||||
|
) -> List[Tuple[Currency, Currency, Decimal]]:
|
||||||
|
results = []
|
||||||
|
ticker_factory = self._get_ticker_factory()
|
||||||
|
|
||||||
|
for asset in target_currencies:
|
||||||
|
exchange_currency = asset.exchange_currency
|
||||||
|
if exchange_currency not in exchange_currencies:
|
||||||
|
continue
|
||||||
|
|
||||||
|
try:
|
||||||
|
history = ticker_factory(asset.code).history(
|
||||||
|
period="5d", interval="1h", auto_adjust=False
|
||||||
|
)
|
||||||
|
|
||||||
|
if history is None or history.empty:
|
||||||
|
logger.warning(
|
||||||
|
"YFinanceMarkets: no history returned for %s", asset.code
|
||||||
|
)
|
||||||
|
continue
|
||||||
|
|
||||||
|
try:
|
||||||
|
latest_close = history["Close"].dropna().iloc[-1]
|
||||||
|
except (IndexError, KeyError, TypeError):
|
||||||
|
logger.warning(
|
||||||
|
"YFinanceMarkets: no close price returned for %s", asset.code
|
||||||
|
)
|
||||||
|
continue
|
||||||
|
|
||||||
|
rate = Decimal(str(latest_close))
|
||||||
|
if not rate.is_finite() or rate <= 0:
|
||||||
|
logger.warning(
|
||||||
|
"YFinanceMarkets: invalid close price %r for %s",
|
||||||
|
latest_close,
|
||||||
|
asset.code,
|
||||||
|
)
|
||||||
|
continue
|
||||||
|
|
||||||
|
results.append((exchange_currency, asset, rate))
|
||||||
|
except Exception as exc:
|
||||||
|
logger.error(
|
||||||
|
"YFinanceMarkets: error fetching %s: %s", asset.code, exc
|
||||||
|
)
|
||||||
|
|
||||||
|
return results
|
||||||
|
|||||||
@@ -0,0 +1,18 @@
|
|||||||
|
# Generated by Django 5.2.10 on 2026-01-10 06:08
|
||||||
|
|
||||||
|
from django.db import migrations, models
|
||||||
|
|
||||||
|
|
||||||
|
class Migration(migrations.Migration):
|
||||||
|
|
||||||
|
dependencies = [
|
||||||
|
('currencies', '0022_currency_is_archived'),
|
||||||
|
]
|
||||||
|
|
||||||
|
operations = [
|
||||||
|
migrations.AddField(
|
||||||
|
model_name='exchangerateservice',
|
||||||
|
name='failure_count',
|
||||||
|
field=models.PositiveIntegerField(default=0),
|
||||||
|
),
|
||||||
|
]
|
||||||
@@ -0,0 +1,30 @@
|
|||||||
|
# Generated by Django 5.2.15 on 2026-07-18 17:26
|
||||||
|
|
||||||
|
from django.db import migrations, models
|
||||||
|
|
||||||
|
|
||||||
|
class Migration(migrations.Migration):
|
||||||
|
|
||||||
|
dependencies = [
|
||||||
|
("currencies", "0023_add_failure_count"),
|
||||||
|
]
|
||||||
|
|
||||||
|
operations = [
|
||||||
|
migrations.AlterField(
|
||||||
|
model_name="exchangerateservice",
|
||||||
|
name="service_type",
|
||||||
|
field=models.CharField(
|
||||||
|
choices=[
|
||||||
|
("coingecko_free", "CoinGecko (Demo/Free)"),
|
||||||
|
("coingecko_pro", "CoinGecko (Pro)"),
|
||||||
|
("transitive", "Transitive (Calculated from Existing Rates)"),
|
||||||
|
("frankfurter", "Frankfurter"),
|
||||||
|
("twelvedata", "TwelveData"),
|
||||||
|
("twelvedatamarkets", "TwelveData Markets"),
|
||||||
|
("yfinance", "Yahoo Finance"),
|
||||||
|
],
|
||||||
|
max_length=255,
|
||||||
|
verbose_name="Service Type",
|
||||||
|
),
|
||||||
|
),
|
||||||
|
]
|
||||||
@@ -105,6 +105,7 @@ class ExchangeRateService(models.Model):
|
|||||||
FRANKFURTER = "frankfurter", "Frankfurter"
|
FRANKFURTER = "frankfurter", "Frankfurter"
|
||||||
TWELVEDATA = "twelvedata", "TwelveData"
|
TWELVEDATA = "twelvedata", "TwelveData"
|
||||||
TWELVEDATA_MARKETS = "twelvedatamarkets", "TwelveData Markets"
|
TWELVEDATA_MARKETS = "twelvedatamarkets", "TwelveData Markets"
|
||||||
|
YFINANCE = "yfinance", "Yahoo Finance"
|
||||||
|
|
||||||
class IntervalType(models.TextChoices):
|
class IntervalType(models.TextChoices):
|
||||||
ON = "on", _("On")
|
ON = "on", _("On")
|
||||||
@@ -136,6 +137,8 @@ class ExchangeRateService(models.Model):
|
|||||||
null=True, blank=True, verbose_name=_("Last Successful Fetch")
|
null=True, blank=True, verbose_name=_("Last Successful Fetch")
|
||||||
)
|
)
|
||||||
|
|
||||||
|
failure_count = models.PositiveIntegerField(default=0)
|
||||||
|
|
||||||
target_currencies = models.ManyToManyField(
|
target_currencies = models.ManyToManyField(
|
||||||
Currency,
|
Currency,
|
||||||
verbose_name=_("Target Currencies"),
|
verbose_name=_("Target Currencies"),
|
||||||
@@ -237,7 +240,7 @@ class ExchangeRateService(models.Model):
|
|||||||
hours = self._parse_hour_ranges(self.fetch_interval)
|
hours = self._parse_hour_ranges(self.fetch_interval)
|
||||||
# Store in normalized format (optional)
|
# Store in normalized format (optional)
|
||||||
self.fetch_interval = ",".join(str(h) for h in sorted(hours))
|
self.fetch_interval = ",".join(str(h) for h in sorted(hours))
|
||||||
except ValueError as e:
|
except ValueError:
|
||||||
raise ValidationError(
|
raise ValidationError(
|
||||||
{
|
{
|
||||||
"fetch_interval": _(
|
"fetch_interval": _(
|
||||||
@@ -248,7 +251,7 @@ class ExchangeRateService(models.Model):
|
|||||||
)
|
)
|
||||||
except ValidationError:
|
except ValidationError:
|
||||||
raise
|
raise
|
||||||
except Exception as e:
|
except Exception:
|
||||||
raise ValidationError(
|
raise ValidationError(
|
||||||
{
|
{
|
||||||
"fetch_interval": _(
|
"fetch_interval": _(
|
||||||
|
|||||||
@@ -0,0 +1 @@
|
|||||||
|
# Tests package for currencies app
|
||||||
@@ -0,0 +1,109 @@
|
|||||||
|
from decimal import Decimal
|
||||||
|
from unittest.mock import patch, MagicMock
|
||||||
|
|
||||||
|
from django.test import TestCase
|
||||||
|
from django.utils import timezone
|
||||||
|
|
||||||
|
from apps.currencies.models import Currency, ExchangeRateService
|
||||||
|
from apps.currencies.exchange_rates.fetcher import ExchangeRateFetcher
|
||||||
|
|
||||||
|
|
||||||
|
class ExchangeRateServiceFailureTrackingTests(TestCase):
|
||||||
|
"""Tests for the failure count tracking functionality."""
|
||||||
|
|
||||||
|
def setUp(self):
|
||||||
|
"""Set up test data."""
|
||||||
|
self.usd = Currency.objects.create(
|
||||||
|
code="USD", name="US Dollar", decimal_places=2, prefix="$ "
|
||||||
|
)
|
||||||
|
self.eur = Currency.objects.create(
|
||||||
|
code="EUR", name="Euro", decimal_places=2, prefix="€ "
|
||||||
|
)
|
||||||
|
self.eur.exchange_currency = self.usd
|
||||||
|
self.eur.save()
|
||||||
|
|
||||||
|
self.service = ExchangeRateService.objects.create(
|
||||||
|
name="Test Service",
|
||||||
|
service_type=ExchangeRateService.ServiceType.FRANKFURTER,
|
||||||
|
is_active=True,
|
||||||
|
)
|
||||||
|
self.service.target_currencies.add(self.eur)
|
||||||
|
|
||||||
|
def test_failure_count_increments_on_provider_error(self):
|
||||||
|
"""Test that failure_count increments when provider raises an exception."""
|
||||||
|
self.assertEqual(self.service.failure_count, 0)
|
||||||
|
|
||||||
|
with patch.object(
|
||||||
|
self.service, "get_provider", side_effect=Exception("API Error")
|
||||||
|
):
|
||||||
|
ExchangeRateFetcher._fetch_service_rates(self.service)
|
||||||
|
|
||||||
|
self.service.refresh_from_db()
|
||||||
|
self.assertEqual(self.service.failure_count, 1)
|
||||||
|
|
||||||
|
def test_failure_count_resets_on_success(self):
|
||||||
|
"""Test that failure_count resets to 0 on successful fetch."""
|
||||||
|
# Set initial failure count
|
||||||
|
self.service.failure_count = 5
|
||||||
|
self.service.save()
|
||||||
|
|
||||||
|
# Mock a successful provider
|
||||||
|
mock_provider = MagicMock()
|
||||||
|
mock_provider.requires_api_key.return_value = False
|
||||||
|
mock_provider.get_rates.return_value = [(self.usd, self.eur, Decimal("0.85"))]
|
||||||
|
mock_provider.rates_inverted = False
|
||||||
|
|
||||||
|
with patch.object(self.service, "get_provider", return_value=mock_provider):
|
||||||
|
ExchangeRateFetcher._fetch_service_rates(self.service)
|
||||||
|
|
||||||
|
self.service.refresh_from_db()
|
||||||
|
self.assertEqual(self.service.failure_count, 0)
|
||||||
|
|
||||||
|
def test_failure_count_accumulates_across_fetches(self):
|
||||||
|
"""Test that failure_count accumulates with consecutive failures."""
|
||||||
|
self.assertEqual(self.service.failure_count, 0)
|
||||||
|
|
||||||
|
with patch.object(
|
||||||
|
self.service, "get_provider", side_effect=Exception("API Error")
|
||||||
|
):
|
||||||
|
ExchangeRateFetcher._fetch_service_rates(self.service)
|
||||||
|
self.service.refresh_from_db()
|
||||||
|
self.assertEqual(self.service.failure_count, 1)
|
||||||
|
|
||||||
|
ExchangeRateFetcher._fetch_service_rates(self.service)
|
||||||
|
self.service.refresh_from_db()
|
||||||
|
self.assertEqual(self.service.failure_count, 2)
|
||||||
|
|
||||||
|
ExchangeRateFetcher._fetch_service_rates(self.service)
|
||||||
|
self.service.refresh_from_db()
|
||||||
|
self.assertEqual(self.service.failure_count, 3)
|
||||||
|
|
||||||
|
def test_last_fetch_not_updated_on_failure(self):
|
||||||
|
"""Test that last_fetch is NOT updated when a failure occurs."""
|
||||||
|
original_last_fetch = self.service.last_fetch
|
||||||
|
self.assertIsNone(original_last_fetch)
|
||||||
|
|
||||||
|
with patch.object(
|
||||||
|
self.service, "get_provider", side_effect=Exception("API Error")
|
||||||
|
):
|
||||||
|
ExchangeRateFetcher._fetch_service_rates(self.service)
|
||||||
|
|
||||||
|
self.service.refresh_from_db()
|
||||||
|
self.assertIsNone(self.service.last_fetch)
|
||||||
|
self.assertEqual(self.service.failure_count, 1)
|
||||||
|
|
||||||
|
def test_last_fetch_updated_on_success(self):
|
||||||
|
"""Test that last_fetch IS updated when fetch succeeds."""
|
||||||
|
self.assertIsNone(self.service.last_fetch)
|
||||||
|
|
||||||
|
mock_provider = MagicMock()
|
||||||
|
mock_provider.requires_api_key.return_value = False
|
||||||
|
mock_provider.get_rates.return_value = [(self.usd, self.eur, Decimal("0.85"))]
|
||||||
|
mock_provider.rates_inverted = False
|
||||||
|
|
||||||
|
with patch.object(self.service, "get_provider", return_value=mock_provider):
|
||||||
|
ExchangeRateFetcher._fetch_service_rates(self.service)
|
||||||
|
|
||||||
|
self.service.refresh_from_db()
|
||||||
|
self.assertIsNotNone(self.service.last_fetch)
|
||||||
|
self.assertEqual(self.service.failure_count, 0)
|
||||||
@@ -0,0 +1,93 @@
|
|||||||
|
from decimal import Decimal
|
||||||
|
from unittest import TestCase
|
||||||
|
|
||||||
|
from apps.currencies.exchange_rates.fetcher import PROVIDER_MAPPING
|
||||||
|
from apps.currencies.exchange_rates.providers import YFinanceMarketsProvider
|
||||||
|
from apps.currencies.models import ExchangeRateService
|
||||||
|
|
||||||
|
|
||||||
|
class _FakeSeries:
|
||||||
|
def __init__(self, values):
|
||||||
|
self._values = values
|
||||||
|
|
||||||
|
def dropna(self):
|
||||||
|
return _FakeSeries([value for value in self._values if value is not None])
|
||||||
|
|
||||||
|
@property
|
||||||
|
def iloc(self):
|
||||||
|
return self
|
||||||
|
|
||||||
|
def __getitem__(self, index):
|
||||||
|
return self._values[index]
|
||||||
|
|
||||||
|
|
||||||
|
class _FakeHistory:
|
||||||
|
def __init__(self, close_values):
|
||||||
|
self._close_values = close_values
|
||||||
|
self.empty = not close_values
|
||||||
|
|
||||||
|
def __getitem__(self, field):
|
||||||
|
if field != "Close":
|
||||||
|
raise KeyError(field)
|
||||||
|
return _FakeSeries(self._close_values)
|
||||||
|
|
||||||
|
|
||||||
|
class _FakeCurrency:
|
||||||
|
def __init__(self, code, exchange_currency=None):
|
||||||
|
self.code = code
|
||||||
|
self.exchange_currency = exchange_currency
|
||||||
|
|
||||||
|
|
||||||
|
class _FakeTicker:
|
||||||
|
def __init__(self, history):
|
||||||
|
self.history_result = history
|
||||||
|
self.history_kwargs = None
|
||||||
|
|
||||||
|
def history(self, **kwargs):
|
||||||
|
self.history_kwargs = kwargs
|
||||||
|
return self.history_result
|
||||||
|
|
||||||
|
|
||||||
|
class YFinanceMarketsProviderTests(TestCase):
|
||||||
|
def setUp(self):
|
||||||
|
self.brl = _FakeCurrency("BRL")
|
||||||
|
self.asset = _FakeCurrency("AAPL", exchange_currency=self.brl)
|
||||||
|
|
||||||
|
def test_returns_latest_hourly_close_using_symbol_verbatim(self):
|
||||||
|
ticker = _FakeTicker(_FakeHistory([36.90, None, 37.42]))
|
||||||
|
requested_symbols = []
|
||||||
|
|
||||||
|
def ticker_factory(symbol):
|
||||||
|
requested_symbols.append(symbol)
|
||||||
|
return ticker
|
||||||
|
|
||||||
|
provider = YFinanceMarketsProvider(ticker_factory=ticker_factory)
|
||||||
|
|
||||||
|
rates = provider.get_rates([self.asset], {self.brl})
|
||||||
|
|
||||||
|
self.assertEqual(rates, [(self.brl, self.asset, Decimal("37.42"))])
|
||||||
|
self.assertEqual(requested_symbols, ["AAPL"])
|
||||||
|
self.assertEqual(
|
||||||
|
ticker.history_kwargs,
|
||||||
|
{"period": "5d", "interval": "1h", "auto_adjust": False},
|
||||||
|
)
|
||||||
|
|
||||||
|
def test_passes_brazilian_symbol_verbatim_and_skips_empty_history(self):
|
||||||
|
self.asset.code = "PETR4.SA"
|
||||||
|
ticker = _FakeTicker(_FakeHistory([]))
|
||||||
|
requested_symbols = []
|
||||||
|
|
||||||
|
provider = YFinanceMarketsProvider(
|
||||||
|
ticker_factory=lambda symbol: requested_symbols.append(symbol)
|
||||||
|
or ticker
|
||||||
|
)
|
||||||
|
|
||||||
|
rates = provider.get_rates([self.asset], {self.brl})
|
||||||
|
|
||||||
|
self.assertEqual(rates, [])
|
||||||
|
self.assertEqual(requested_symbols, ["PETR4.SA"])
|
||||||
|
|
||||||
|
def test_is_registered_without_an_api_key(self):
|
||||||
|
self.assertFalse(YFinanceMarketsProvider.requires_api_key())
|
||||||
|
self.assertIs(PROVIDER_MAPPING["yfinance"], YFinanceMarketsProvider)
|
||||||
|
self.assertEqual(ExchangeRateService.ServiceType.YFINANCE, "yfinance")
|
||||||
@@ -23,7 +23,7 @@ def currencies_index(request):
|
|||||||
@login_required
|
@login_required
|
||||||
@require_http_methods(["GET"])
|
@require_http_methods(["GET"])
|
||||||
def currencies_list(request):
|
def currencies_list(request):
|
||||||
currencies = Currency.objects.all().order_by("id")
|
currencies = Currency.objects.all().order_by("name")
|
||||||
return render(
|
return render(
|
||||||
request,
|
request,
|
||||||
"currencies/fragments/list.html",
|
"currencies/fragments/list.html",
|
||||||
|
|||||||
@@ -1,3 +0,0 @@
|
|||||||
from django.test import TestCase
|
|
||||||
|
|
||||||
# Create your tests here.
|
|
||||||
@@ -0,0 +1,248 @@
|
|||||||
|
"""Object-level authorization tests for the DCA views.
|
||||||
|
|
||||||
|
DCAEntry has an unscoped default manager, so filtering an entry by
|
||||||
|
``strategy__id`` alone reached entries belonging to strategies the caller could
|
||||||
|
not see at all -- a strictly wider hole than the one reported for rules in
|
||||||
|
GHSA-83g9-vjqf-2j5q, since it needs no public or shared strategy.
|
||||||
|
"""
|
||||||
|
|
||||||
|
from datetime import date
|
||||||
|
from decimal import Decimal
|
||||||
|
|
||||||
|
from django.contrib.auth import get_user_model
|
||||||
|
from django.test import TestCase, override_settings
|
||||||
|
from django.urls import reverse
|
||||||
|
|
||||||
|
from apps.currencies.models import Currency
|
||||||
|
from apps.dca.models import DCAEntry, DCAStrategy
|
||||||
|
|
||||||
|
HTMX = {"HTTP_HX_REQUEST": "true"}
|
||||||
|
|
||||||
|
|
||||||
|
@override_settings(
|
||||||
|
STORAGES={
|
||||||
|
"default": {"BACKEND": "django.core.files.storage.FileSystemStorage"},
|
||||||
|
"staticfiles": {
|
||||||
|
"BACKEND": "django.contrib.staticfiles.storage.StaticFilesStorage"
|
||||||
|
},
|
||||||
|
},
|
||||||
|
WHITENOISE_AUTOREFRESH=True,
|
||||||
|
DEMO=False,
|
||||||
|
)
|
||||||
|
class DCAObjectPermissionTests(TestCase):
|
||||||
|
def setUp(self):
|
||||||
|
User = get_user_model()
|
||||||
|
self.owner = User.objects.create_user(
|
||||||
|
email="owner@test.com", password="testpass123"
|
||||||
|
)
|
||||||
|
self.shared_user = User.objects.create_user(
|
||||||
|
email="shared@test.com", password="testpass123"
|
||||||
|
)
|
||||||
|
self.stranger = User.objects.create_user(
|
||||||
|
email="stranger@test.com", password="testpass123"
|
||||||
|
)
|
||||||
|
self.currency = Currency.objects.create(
|
||||||
|
code="USD", name="US Dollar", decimal_places=2
|
||||||
|
)
|
||||||
|
|
||||||
|
self.private_strategy = self._strategy("Private", visibility="private")
|
||||||
|
self.public_strategy = self._strategy("Public", visibility="public")
|
||||||
|
self.shared_strategy = self._strategy("Shared", visibility="private")
|
||||||
|
self.shared_strategy.shared_with.add(self.shared_user)
|
||||||
|
|
||||||
|
self.private_entry = self._entry(self.private_strategy)
|
||||||
|
self.public_entry = self._entry(self.public_strategy)
|
||||||
|
self.shared_entry = self._entry(self.shared_strategy)
|
||||||
|
|
||||||
|
def _strategy(self, name, visibility):
|
||||||
|
return DCAStrategy.all_objects.create(
|
||||||
|
name=name,
|
||||||
|
owner=self.owner,
|
||||||
|
visibility=visibility,
|
||||||
|
target_currency=self.currency,
|
||||||
|
payment_currency=self.currency,
|
||||||
|
)
|
||||||
|
|
||||||
|
def _entry(self, strategy):
|
||||||
|
return DCAEntry.objects.create(
|
||||||
|
strategy=strategy,
|
||||||
|
date=date(2025, 1, 1),
|
||||||
|
amount_paid=Decimal("100"),
|
||||||
|
amount_received=Decimal("1"),
|
||||||
|
)
|
||||||
|
|
||||||
|
# ------------------------------------------------------------------
|
||||||
|
# entries on an invisible strategy: must not even confirm they exist
|
||||||
|
# ------------------------------------------------------------------
|
||||||
|
def test_stranger_cannot_delete_entry_on_private_strategy(self):
|
||||||
|
self.client.force_login(self.stranger)
|
||||||
|
|
||||||
|
response = self.client.delete(
|
||||||
|
reverse(
|
||||||
|
"dca_entry_delete",
|
||||||
|
kwargs={
|
||||||
|
"strategy_id": self.private_strategy.id,
|
||||||
|
"entry_id": self.private_entry.id,
|
||||||
|
},
|
||||||
|
),
|
||||||
|
**HTMX,
|
||||||
|
)
|
||||||
|
|
||||||
|
self.assertEqual(response.status_code, 404)
|
||||||
|
self.assertTrue(DCAEntry.objects.filter(pk=self.private_entry.pk).exists())
|
||||||
|
|
||||||
|
def test_stranger_cannot_open_entry_edit_form_on_private_strategy(self):
|
||||||
|
self.client.force_login(self.stranger)
|
||||||
|
|
||||||
|
response = self.client.get(
|
||||||
|
reverse(
|
||||||
|
"dca_entry_edit",
|
||||||
|
kwargs={
|
||||||
|
"strategy_id": self.private_strategy.id,
|
||||||
|
"entry_id": self.private_entry.id,
|
||||||
|
},
|
||||||
|
),
|
||||||
|
**HTMX,
|
||||||
|
)
|
||||||
|
|
||||||
|
self.assertEqual(response.status_code, 404)
|
||||||
|
|
||||||
|
def test_stranger_cannot_edit_entry_on_private_strategy(self):
|
||||||
|
self.client.force_login(self.stranger)
|
||||||
|
|
||||||
|
response = self.client.post(
|
||||||
|
reverse(
|
||||||
|
"dca_entry_edit",
|
||||||
|
kwargs={
|
||||||
|
"strategy_id": self.private_strategy.id,
|
||||||
|
"entry_id": self.private_entry.id,
|
||||||
|
},
|
||||||
|
),
|
||||||
|
data={
|
||||||
|
"date": "2030-01-01",
|
||||||
|
"amount_paid": "999",
|
||||||
|
"amount_received": "999",
|
||||||
|
},
|
||||||
|
**HTMX,
|
||||||
|
)
|
||||||
|
|
||||||
|
self.assertEqual(response.status_code, 404)
|
||||||
|
self.private_entry.refresh_from_db()
|
||||||
|
self.assertEqual(self.private_entry.amount_paid, Decimal("100"))
|
||||||
|
|
||||||
|
# ------------------------------------------------------------------
|
||||||
|
# entries on a visible-but-unowned strategy: 403, not 404
|
||||||
|
# ------------------------------------------------------------------
|
||||||
|
def test_stranger_cannot_delete_entry_on_public_strategy(self):
|
||||||
|
self.client.force_login(self.stranger)
|
||||||
|
|
||||||
|
response = self.client.delete(
|
||||||
|
reverse(
|
||||||
|
"dca_entry_delete",
|
||||||
|
kwargs={
|
||||||
|
"strategy_id": self.public_strategy.id,
|
||||||
|
"entry_id": self.public_entry.id,
|
||||||
|
},
|
||||||
|
),
|
||||||
|
**HTMX,
|
||||||
|
)
|
||||||
|
|
||||||
|
self.assertEqual(response.status_code, 403)
|
||||||
|
self.assertTrue(DCAEntry.objects.filter(pk=self.public_entry.pk).exists())
|
||||||
|
|
||||||
|
def test_shared_user_cannot_add_entry_to_shared_strategy(self):
|
||||||
|
self.client.force_login(self.shared_user)
|
||||||
|
|
||||||
|
response = self.client.post(
|
||||||
|
reverse("dca_entry_add", kwargs={"strategy_id": self.shared_strategy.id}),
|
||||||
|
data={
|
||||||
|
"date": "2030-01-01",
|
||||||
|
"amount_paid": "5",
|
||||||
|
"amount_received": "5",
|
||||||
|
},
|
||||||
|
**HTMX,
|
||||||
|
)
|
||||||
|
|
||||||
|
self.assertEqual(response.status_code, 403)
|
||||||
|
self.assertEqual(self.shared_strategy.entries.count(), 1)
|
||||||
|
|
||||||
|
def test_owner_can_still_manage_own_entries(self):
|
||||||
|
self.client.force_login(self.owner)
|
||||||
|
|
||||||
|
response = self.client.delete(
|
||||||
|
reverse(
|
||||||
|
"dca_entry_delete",
|
||||||
|
kwargs={
|
||||||
|
"strategy_id": self.private_strategy.id,
|
||||||
|
"entry_id": self.private_entry.id,
|
||||||
|
},
|
||||||
|
),
|
||||||
|
**HTMX,
|
||||||
|
)
|
||||||
|
|
||||||
|
self.assertEqual(response.status_code, 204)
|
||||||
|
self.assertFalse(DCAEntry.objects.filter(pk=self.private_entry.pk).exists())
|
||||||
|
|
||||||
|
# ------------------------------------------------------------------
|
||||||
|
# reads stay open to shared users
|
||||||
|
# ------------------------------------------------------------------
|
||||||
|
def test_shared_user_can_still_view_strategy_detail(self):
|
||||||
|
self.client.force_login(self.shared_user)
|
||||||
|
|
||||||
|
response = self.client.get(
|
||||||
|
reverse(
|
||||||
|
"dca_strategy_detail", kwargs={"strategy_id": self.shared_strategy.id}
|
||||||
|
),
|
||||||
|
**HTMX,
|
||||||
|
)
|
||||||
|
|
||||||
|
self.assertEqual(response.status_code, 200)
|
||||||
|
|
||||||
|
# ------------------------------------------------------------------
|
||||||
|
# strategy delete: the same inversion the rules views had
|
||||||
|
# ------------------------------------------------------------------
|
||||||
|
def test_stranger_cannot_delete_public_strategy(self):
|
||||||
|
self.client.force_login(self.stranger)
|
||||||
|
|
||||||
|
response = self.client.delete(
|
||||||
|
reverse(
|
||||||
|
"dca_strategy_delete", kwargs={"strategy_id": self.public_strategy.id}
|
||||||
|
),
|
||||||
|
**HTMX,
|
||||||
|
)
|
||||||
|
|
||||||
|
self.assertEqual(response.status_code, 403)
|
||||||
|
self.assertTrue(
|
||||||
|
DCAStrategy.all_objects.filter(pk=self.public_strategy.pk).exists()
|
||||||
|
)
|
||||||
|
|
||||||
|
def test_shared_user_deleting_only_revokes_their_own_access(self):
|
||||||
|
self.client.force_login(self.shared_user)
|
||||||
|
|
||||||
|
response = self.client.delete(
|
||||||
|
reverse(
|
||||||
|
"dca_strategy_delete", kwargs={"strategy_id": self.shared_strategy.id}
|
||||||
|
),
|
||||||
|
**HTMX,
|
||||||
|
)
|
||||||
|
|
||||||
|
self.assertEqual(response.status_code, 204)
|
||||||
|
self.assertTrue(
|
||||||
|
DCAStrategy.all_objects.filter(pk=self.shared_strategy.pk).exists()
|
||||||
|
)
|
||||||
|
self.assertNotIn(self.shared_user, self.shared_strategy.shared_with.all())
|
||||||
|
|
||||||
|
def test_owner_can_delete_own_strategy(self):
|
||||||
|
self.client.force_login(self.owner)
|
||||||
|
|
||||||
|
response = self.client.delete(
|
||||||
|
reverse(
|
||||||
|
"dca_strategy_delete", kwargs={"strategy_id": self.public_strategy.id}
|
||||||
|
),
|
||||||
|
**HTMX,
|
||||||
|
)
|
||||||
|
|
||||||
|
self.assertEqual(response.status_code, 204)
|
||||||
|
self.assertFalse(
|
||||||
|
DCAStrategy.all_objects.filter(pk=self.public_strategy.pk).exists()
|
||||||
|
)
|
||||||
+51
-39
@@ -1,14 +1,19 @@
|
|||||||
# apps/dca_tracker/views.py
|
|
||||||
from django.contrib import messages
|
from django.contrib import messages
|
||||||
from django.contrib.auth.decorators import login_required
|
from django.contrib.auth.decorators import login_required
|
||||||
|
from django.core.exceptions import PermissionDenied
|
||||||
from django.db.models import Sum, Avg
|
from django.db.models import Sum, Avg
|
||||||
from django.db.models.functions import TruncMonth
|
from django.db.models.functions import TruncMonth
|
||||||
from django.http import HttpResponse
|
from django.http import HttpResponse
|
||||||
from django.shortcuts import render, get_object_or_404
|
from django.shortcuts import render
|
||||||
from django.utils.translation import gettext_lazy as _
|
from django.utils.translation import gettext_lazy as _
|
||||||
from django.views.decorators.http import require_http_methods
|
from django.views.decorators.http import require_http_methods
|
||||||
|
|
||||||
from apps.common.decorators.htmx import only_htmx
|
from apps.common.decorators.htmx import only_htmx
|
||||||
|
from apps.common.functions.permissions import (
|
||||||
|
EDIT,
|
||||||
|
READ,
|
||||||
|
get_shared_object_or_error,
|
||||||
|
)
|
||||||
from apps.dca.forms import DCAEntryForm, DCAStrategyForm
|
from apps.dca.forms import DCAEntryForm, DCAStrategyForm
|
||||||
from apps.dca.models import DCAStrategy, DCAEntry
|
from apps.dca.models import DCAStrategy, DCAEntry
|
||||||
from apps.common.models import SharedObject
|
from apps.common.models import SharedObject
|
||||||
@@ -23,7 +28,7 @@ def strategy_index(request):
|
|||||||
@only_htmx
|
@only_htmx
|
||||||
@login_required
|
@login_required
|
||||||
def strategy_list(request):
|
def strategy_list(request):
|
||||||
strategies = DCAStrategy.objects.all().order_by("created_at")
|
strategies = DCAStrategy.objects.all().order_by("name")
|
||||||
return render(
|
return render(
|
||||||
request, "dca/fragments/strategy/list.html", {"strategies": strategies}
|
request, "dca/fragments/strategy/list.html", {"strategies": strategies}
|
||||||
)
|
)
|
||||||
@@ -57,17 +62,9 @@ def strategy_add(request):
|
|||||||
@only_htmx
|
@only_htmx
|
||||||
@login_required
|
@login_required
|
||||||
def strategy_edit(request, strategy_id):
|
def strategy_edit(request, strategy_id):
|
||||||
dca_strategy = get_object_or_404(DCAStrategy, id=strategy_id)
|
dca_strategy = get_shared_object_or_error(
|
||||||
|
DCAStrategy, request, id=strategy_id, level=EDIT
|
||||||
if dca_strategy.owner and dca_strategy.owner != request.user:
|
)
|
||||||
messages.error(request, _("Only the owner can edit this"))
|
|
||||||
|
|
||||||
return HttpResponse(
|
|
||||||
status=204,
|
|
||||||
headers={
|
|
||||||
"HX-Trigger": "updated, hide_offcanvas",
|
|
||||||
},
|
|
||||||
)
|
|
||||||
|
|
||||||
if request.method == "POST":
|
if request.method == "POST":
|
||||||
form = DCAStrategyForm(request.POST, instance=dca_strategy)
|
form = DCAStrategyForm(request.POST, instance=dca_strategy)
|
||||||
@@ -95,17 +92,20 @@ def strategy_edit(request, strategy_id):
|
|||||||
@login_required
|
@login_required
|
||||||
@require_http_methods(["DELETE"])
|
@require_http_methods(["DELETE"])
|
||||||
def strategy_delete(request, strategy_id):
|
def strategy_delete(request, strategy_id):
|
||||||
dca_strategy = get_object_or_404(DCAStrategy, id=strategy_id)
|
dca_strategy = get_shared_object_or_error(
|
||||||
|
DCAStrategy, request, id=strategy_id, level=READ
|
||||||
|
)
|
||||||
|
|
||||||
if (
|
if dca_strategy.is_editable_by(request.user):
|
||||||
dca_strategy.owner != request.user
|
dca_strategy.delete()
|
||||||
and request.user in dca_strategy.shared_with.all()
|
messages.success(request, _("DCA strategy deleted successfully"))
|
||||||
):
|
elif dca_strategy.shared_with.filter(pk=request.user.pk).exists():
|
||||||
|
# Someone else's object shared with us: we can drop our own access
|
||||||
|
# to it, but never delete it.
|
||||||
dca_strategy.shared_with.remove(request.user)
|
dca_strategy.shared_with.remove(request.user)
|
||||||
messages.success(request, _("Item no longer shared with you"))
|
messages.success(request, _("Item no longer shared with you"))
|
||||||
else:
|
else:
|
||||||
dca_strategy.delete()
|
raise PermissionDenied
|
||||||
messages.success(request, _("DCA strategy deleted successfully"))
|
|
||||||
|
|
||||||
return HttpResponse(
|
return HttpResponse(
|
||||||
status=204,
|
status=204,
|
||||||
@@ -119,7 +119,9 @@ def strategy_delete(request, strategy_id):
|
|||||||
@login_required
|
@login_required
|
||||||
@require_http_methods(["GET"])
|
@require_http_methods(["GET"])
|
||||||
def strategy_take_ownership(request, strategy_id):
|
def strategy_take_ownership(request, strategy_id):
|
||||||
dca_strategy = get_object_or_404(DCAStrategy, id=strategy_id)
|
dca_strategy = get_shared_object_or_error(
|
||||||
|
DCAStrategy, request, id=strategy_id, level=EDIT
|
||||||
|
)
|
||||||
|
|
||||||
if not dca_strategy.owner:
|
if not dca_strategy.owner:
|
||||||
dca_strategy.owner = request.user
|
dca_strategy.owner = request.user
|
||||||
@@ -140,17 +142,7 @@ def strategy_take_ownership(request, strategy_id):
|
|||||||
@login_required
|
@login_required
|
||||||
@require_http_methods(["GET", "POST"])
|
@require_http_methods(["GET", "POST"])
|
||||||
def strategy_share(request, pk):
|
def strategy_share(request, pk):
|
||||||
obj = get_object_or_404(DCAStrategy, id=pk)
|
obj = get_shared_object_or_error(DCAStrategy, request, id=pk, level=EDIT)
|
||||||
|
|
||||||
if obj.owner and obj.owner != request.user:
|
|
||||||
messages.error(request, _("Only the owner can edit this"))
|
|
||||||
|
|
||||||
return HttpResponse(
|
|
||||||
status=204,
|
|
||||||
headers={
|
|
||||||
"HX-Trigger": "updated, hide_offcanvas",
|
|
||||||
},
|
|
||||||
)
|
|
||||||
|
|
||||||
if request.method == "POST":
|
if request.method == "POST":
|
||||||
form = SharedObjectForm(request.POST, instance=obj, user=request.user)
|
form = SharedObjectForm(request.POST, instance=obj, user=request.user)
|
||||||
@@ -176,7 +168,9 @@ def strategy_share(request, pk):
|
|||||||
|
|
||||||
@login_required
|
@login_required
|
||||||
def strategy_detail_index(request, strategy_id):
|
def strategy_detail_index(request, strategy_id):
|
||||||
strategy = get_object_or_404(DCAStrategy, id=strategy_id)
|
strategy = get_shared_object_or_error(
|
||||||
|
DCAStrategy, request, id=strategy_id, level=READ
|
||||||
|
)
|
||||||
|
|
||||||
return render(
|
return render(
|
||||||
request,
|
request,
|
||||||
@@ -188,7 +182,9 @@ def strategy_detail_index(request, strategy_id):
|
|||||||
@only_htmx
|
@only_htmx
|
||||||
@login_required
|
@login_required
|
||||||
def strategy_detail(request, strategy_id):
|
def strategy_detail(request, strategy_id):
|
||||||
strategy = get_object_or_404(DCAStrategy, id=strategy_id)
|
strategy = get_shared_object_or_error(
|
||||||
|
DCAStrategy, request, id=strategy_id, level=READ
|
||||||
|
)
|
||||||
entries = strategy.entries.all()
|
entries = strategy.entries.all()
|
||||||
|
|
||||||
# Calculate monthly aggregates
|
# Calculate monthly aggregates
|
||||||
@@ -230,11 +226,13 @@ def strategy_detail(request, strategy_id):
|
|||||||
@only_htmx
|
@only_htmx
|
||||||
@login_required
|
@login_required
|
||||||
def strategy_entry_add(request, strategy_id):
|
def strategy_entry_add(request, strategy_id):
|
||||||
strategy = get_object_or_404(DCAStrategy, id=strategy_id)
|
strategy = get_shared_object_or_error(
|
||||||
|
DCAStrategy, request, id=strategy_id, level=EDIT
|
||||||
|
)
|
||||||
if request.method == "POST":
|
if request.method == "POST":
|
||||||
form = DCAEntryForm(request.POST, strategy=strategy)
|
form = DCAEntryForm(request.POST, strategy=strategy)
|
||||||
if form.is_valid():
|
if form.is_valid():
|
||||||
entry = form.save()
|
form.save()
|
||||||
messages.success(request, _("Entry added successfully"))
|
messages.success(request, _("Entry added successfully"))
|
||||||
|
|
||||||
return HttpResponse(
|
return HttpResponse(
|
||||||
@@ -256,7 +254,14 @@ def strategy_entry_add(request, strategy_id):
|
|||||||
@only_htmx
|
@only_htmx
|
||||||
@login_required
|
@login_required
|
||||||
def strategy_entry_edit(request, strategy_id, entry_id):
|
def strategy_entry_edit(request, strategy_id, entry_id):
|
||||||
dca_entry = get_object_or_404(DCAEntry, id=entry_id, strategy__id=strategy_id)
|
dca_entry = get_shared_object_or_error(
|
||||||
|
DCAEntry,
|
||||||
|
request,
|
||||||
|
id=entry_id,
|
||||||
|
strategy__id=strategy_id,
|
||||||
|
level=EDIT,
|
||||||
|
via="strategy",
|
||||||
|
)
|
||||||
|
|
||||||
if request.method == "POST":
|
if request.method == "POST":
|
||||||
form = DCAEntryForm(request.POST, instance=dca_entry)
|
form = DCAEntryForm(request.POST, instance=dca_entry)
|
||||||
@@ -284,7 +289,14 @@ def strategy_entry_edit(request, strategy_id, entry_id):
|
|||||||
@login_required
|
@login_required
|
||||||
@require_http_methods(["DELETE"])
|
@require_http_methods(["DELETE"])
|
||||||
def strategy_entry_delete(request, entry_id, strategy_id):
|
def strategy_entry_delete(request, entry_id, strategy_id):
|
||||||
dca_entry = get_object_or_404(DCAEntry, id=entry_id, strategy__id=strategy_id)
|
dca_entry = get_shared_object_or_error(
|
||||||
|
DCAEntry,
|
||||||
|
request,
|
||||||
|
id=entry_id,
|
||||||
|
strategy__id=strategy_id,
|
||||||
|
level=EDIT,
|
||||||
|
via="strategy",
|
||||||
|
)
|
||||||
|
|
||||||
dca_entry.delete()
|
dca_entry.delete()
|
||||||
|
|
||||||
|
|||||||
@@ -1,8 +1,10 @@
|
|||||||
from import_export import fields, resources
|
from import_export import fields, resources
|
||||||
from import_export.widgets import ForeignKeyWidget
|
|
||||||
|
|
||||||
from apps.accounts.models import Account
|
from apps.accounts.models import Account
|
||||||
from apps.export_app.widgets.foreign_key import AutoCreateForeignKeyWidget
|
from apps.export_app.widgets.foreign_key import (
|
||||||
|
AllObjectsForeignKeyWidget,
|
||||||
|
AutoCreateForeignKeyWidget,
|
||||||
|
)
|
||||||
from apps.export_app.widgets.many_to_many import AutoCreateManyToManyWidget
|
from apps.export_app.widgets.many_to_many import AutoCreateManyToManyWidget
|
||||||
from apps.export_app.widgets.string import EmptyStringToNoneField
|
from apps.export_app.widgets.string import EmptyStringToNoneField
|
||||||
from apps.transactions.models import (
|
from apps.transactions.models import (
|
||||||
@@ -20,7 +22,7 @@ class TransactionResource(resources.ModelResource):
|
|||||||
account = fields.Field(
|
account = fields.Field(
|
||||||
attribute="account",
|
attribute="account",
|
||||||
column_name="account",
|
column_name="account",
|
||||||
widget=ForeignKeyWidget(Account, "name"),
|
widget=AllObjectsForeignKeyWidget(Account, "name"),
|
||||||
)
|
)
|
||||||
|
|
||||||
category = fields.Field(
|
category = fields.Field(
|
||||||
@@ -86,7 +88,7 @@ class RecurringTransactionResource(resources.ModelResource):
|
|||||||
account = fields.Field(
|
account = fields.Field(
|
||||||
attribute="account",
|
attribute="account",
|
||||||
column_name="account",
|
column_name="account",
|
||||||
widget=ForeignKeyWidget(Account, "name"),
|
widget=AllObjectsForeignKeyWidget(Account, "name"),
|
||||||
)
|
)
|
||||||
|
|
||||||
category = fields.Field(
|
category = fields.Field(
|
||||||
@@ -119,12 +121,16 @@ class RecurringTransactionResource(resources.ModelResource):
|
|||||||
def get_queryset(self):
|
def get_queryset(self):
|
||||||
return RecurringTransaction.all_objects.all()
|
return RecurringTransaction.all_objects.all()
|
||||||
|
|
||||||
|
def dehydrate_account_owner(self, obj):
|
||||||
|
"""Export the account's owner ID for proper import matching."""
|
||||||
|
return obj.account.owner_id if obj.account else None
|
||||||
|
|
||||||
|
|
||||||
class InstallmentPlanResource(resources.ModelResource):
|
class InstallmentPlanResource(resources.ModelResource):
|
||||||
account = fields.Field(
|
account = fields.Field(
|
||||||
attribute="account",
|
attribute="account",
|
||||||
column_name="account",
|
column_name="account",
|
||||||
widget=ForeignKeyWidget(Account, "name"),
|
widget=AllObjectsForeignKeyWidget(Account, "name"),
|
||||||
)
|
)
|
||||||
|
|
||||||
category = fields.Field(
|
category = fields.Field(
|
||||||
@@ -156,3 +162,7 @@ class InstallmentPlanResource(resources.ModelResource):
|
|||||||
|
|
||||||
def get_queryset(self):
|
def get_queryset(self):
|
||||||
return InstallmentPlan.all_objects.all()
|
return InstallmentPlan.all_objects.all()
|
||||||
|
|
||||||
|
def dehydrate_account_owner(self, obj):
|
||||||
|
"""Export the account's owner ID for proper import matching."""
|
||||||
|
return obj.account.owner_id if obj.account else None
|
||||||
|
|||||||
@@ -1,6 +1,60 @@
|
|||||||
from import_export.widgets import ForeignKeyWidget
|
from import_export.widgets import ForeignKeyWidget
|
||||||
|
|
||||||
|
|
||||||
|
class AllObjectsForeignKeyWidget(ForeignKeyWidget):
|
||||||
|
"""
|
||||||
|
ForeignKeyWidget that uses 'all_objects' manager for lookups,
|
||||||
|
bypassing user-filtered managers like SharedObjectManager.
|
||||||
|
Also filters by owner if available in the row data.
|
||||||
|
"""
|
||||||
|
|
||||||
|
def get_queryset(self, value, row, *args, **kwargs):
|
||||||
|
# Use all_objects manager if available, otherwise fall back to default
|
||||||
|
if hasattr(self.model, "all_objects"):
|
||||||
|
qs = self.model.all_objects.all()
|
||||||
|
# Filter by owner if the row has an owner field and the model has owner
|
||||||
|
if row:
|
||||||
|
# Check for direct owner field first
|
||||||
|
owner_id = row.get("owner") if "owner" in row else None
|
||||||
|
# Fall back to account_owner for models like InstallmentPlan
|
||||||
|
if not owner_id and "account_owner" in row:
|
||||||
|
owner_id = row.get("account_owner")
|
||||||
|
# If still no owner, try to get it from the existing record's account
|
||||||
|
# This handles backward compatibility with older exports
|
||||||
|
if not owner_id and "id" in row and row.get("id"):
|
||||||
|
try:
|
||||||
|
# Try to find the existing record and get owner from its account
|
||||||
|
from apps.transactions.models import (
|
||||||
|
InstallmentPlan,
|
||||||
|
RecurringTransaction,
|
||||||
|
)
|
||||||
|
|
||||||
|
record_id = row.get("id")
|
||||||
|
# Try to find the existing InstallmentPlan or RecurringTransaction
|
||||||
|
for model_class in [InstallmentPlan, RecurringTransaction]:
|
||||||
|
try:
|
||||||
|
existing = model_class.all_objects.get(id=record_id)
|
||||||
|
if existing.account:
|
||||||
|
owner_id = existing.account.owner_id
|
||||||
|
break
|
||||||
|
except model_class.DoesNotExist:
|
||||||
|
continue
|
||||||
|
except Exception:
|
||||||
|
pass
|
||||||
|
# Final fallback: use the current logged-in user
|
||||||
|
# This handles restoring to a fresh database with older exports
|
||||||
|
if not owner_id:
|
||||||
|
from apps.common.middleware.thread_local import get_current_user
|
||||||
|
|
||||||
|
user = get_current_user()
|
||||||
|
if user and user.is_authenticated:
|
||||||
|
owner_id = user.id
|
||||||
|
if owner_id:
|
||||||
|
qs = qs.filter(owner_id=owner_id)
|
||||||
|
return qs
|
||||||
|
return super().get_queryset(value, row, *args, **kwargs)
|
||||||
|
|
||||||
|
|
||||||
class AutoCreateForeignKeyWidget(ForeignKeyWidget):
|
class AutoCreateForeignKeyWidget(ForeignKeyWidget):
|
||||||
def clean(self, value, row=None, *args, **kwargs):
|
def clean(self, value, row=None, *args, **kwargs):
|
||||||
if value:
|
if value:
|
||||||
|
|||||||
@@ -106,6 +106,17 @@ class ExcelImportSettings(BaseModel):
|
|||||||
sheets: list[str] | str = "*"
|
sheets: list[str] | str = "*"
|
||||||
|
|
||||||
|
|
||||||
|
class QIFImportSettings(BaseModel):
|
||||||
|
skip_errors: bool = Field(
|
||||||
|
default=False,
|
||||||
|
description="If True, errors during import will be logged and skipped",
|
||||||
|
)
|
||||||
|
file_type: Literal["qif"] = "qif"
|
||||||
|
importing: Literal["transactions"] = "transactions"
|
||||||
|
encoding: str = Field(default="utf-8", description="File encoding")
|
||||||
|
date_format: str = Field(..., description="Date format (e.g. %d/%m/%Y)")
|
||||||
|
|
||||||
|
|
||||||
class ColumnMapping(BaseModel):
|
class ColumnMapping(BaseModel):
|
||||||
source: Optional[str] | Optional[list[str]] = Field(
|
source: Optional[str] | Optional[list[str]] = Field(
|
||||||
default=None,
|
default=None,
|
||||||
@@ -342,7 +353,7 @@ class CurrencyExchangeMapping(ColumnMapping):
|
|||||||
|
|
||||||
|
|
||||||
class ImportProfileSchema(BaseModel):
|
class ImportProfileSchema(BaseModel):
|
||||||
settings: CSVImportSettings | ExcelImportSettings
|
settings: CSVImportSettings | ExcelImportSettings | QIFImportSettings
|
||||||
mapping: Dict[
|
mapping: Dict[
|
||||||
str,
|
str,
|
||||||
TransactionAccountMapping
|
TransactionAccountMapping
|
||||||
|
|||||||
@@ -3,6 +3,8 @@ import hashlib
|
|||||||
import logging
|
import logging
|
||||||
import os
|
import os
|
||||||
import re
|
import re
|
||||||
|
import zipfile
|
||||||
|
from django.db import transaction
|
||||||
from datetime import datetime, date
|
from datetime import datetime, date
|
||||||
from decimal import Decimal, InvalidOperation
|
from decimal import Decimal, InvalidOperation
|
||||||
from typing import Dict, Any, Literal, Union
|
from typing import Dict, Any, Literal, Union
|
||||||
@@ -11,6 +13,7 @@ import openpyxl
|
|||||||
import xlrd
|
import xlrd
|
||||||
import yaml
|
import yaml
|
||||||
from cachalot.api import cachalot_disabled
|
from cachalot.api import cachalot_disabled
|
||||||
|
from django.core.exceptions import FieldDoesNotExist
|
||||||
from django.utils import timezone
|
from django.utils import timezone
|
||||||
from openpyxl.utils.exceptions import InvalidFileException
|
from openpyxl.utils.exceptions import InvalidFileException
|
||||||
|
|
||||||
@@ -363,7 +366,7 @@ class ImportService:
|
|||||||
try:
|
try:
|
||||||
if entities_mapping:
|
if entities_mapping:
|
||||||
if entities_mapping.type == "id":
|
if entities_mapping.type == "id":
|
||||||
entity = TransactionTag.objects.filter(
|
entity = TransactionEntity.objects.filter(
|
||||||
id=entity_name
|
id=entity_name
|
||||||
).first()
|
).first()
|
||||||
else: # name
|
else: # name
|
||||||
@@ -459,12 +462,13 @@ class ImportService:
|
|||||||
# Build query conditions for each field in the rule
|
# Build query conditions for each field in the rule
|
||||||
for field in rule.fields:
|
for field in rule.fields:
|
||||||
if field in transaction_data:
|
if field in transaction_data:
|
||||||
if rule.match_type == "strict":
|
value = transaction_data[field]
|
||||||
query = query.filter(**{field: transaction_data[field]})
|
query = self._apply_deduplication_filter(
|
||||||
else: # lax matching
|
query=query,
|
||||||
query = query.filter(
|
field=field,
|
||||||
**{f"{field}__iexact": transaction_data[field]}
|
value=value,
|
||||||
)
|
match_type=rule.match_type,
|
||||||
|
)
|
||||||
|
|
||||||
# If we found any matching transaction, it's a duplicate
|
# If we found any matching transaction, it's a duplicate
|
||||||
if query.exists():
|
if query.exists():
|
||||||
@@ -472,6 +476,71 @@ class ImportService:
|
|||||||
|
|
||||||
return False
|
return False
|
||||||
|
|
||||||
|
@staticmethod
|
||||||
|
def _is_int_like(value: Any) -> bool:
|
||||||
|
try:
|
||||||
|
int(value)
|
||||||
|
except (TypeError, ValueError):
|
||||||
|
return False
|
||||||
|
return True
|
||||||
|
|
||||||
|
def _apply_deduplication_filter(
|
||||||
|
self,
|
||||||
|
query,
|
||||||
|
field: str,
|
||||||
|
value: Any,
|
||||||
|
match_type: Literal["lax", "strict"],
|
||||||
|
):
|
||||||
|
if isinstance(value, list):
|
||||||
|
return self._apply_list_deduplication_filter(
|
||||||
|
query=query,
|
||||||
|
field=field,
|
||||||
|
values=value,
|
||||||
|
match_type=match_type,
|
||||||
|
)
|
||||||
|
|
||||||
|
# Use __iexact only for string fields; non-string types
|
||||||
|
# (date, Decimal, bool, int, etc.) don't support UPPER()
|
||||||
|
if match_type == "strict" or not isinstance(value, str):
|
||||||
|
return query.filter(**{field: value})
|
||||||
|
|
||||||
|
return query.filter(**{f"{field}__iexact": value})
|
||||||
|
|
||||||
|
def _apply_list_deduplication_filter(
|
||||||
|
self,
|
||||||
|
query,
|
||||||
|
field: str,
|
||||||
|
values: list[Any],
|
||||||
|
match_type: Literal["lax", "strict"],
|
||||||
|
):
|
||||||
|
clean_values = [v for v in values if v not in (None, "")]
|
||||||
|
if not clean_values:
|
||||||
|
return query
|
||||||
|
|
||||||
|
try:
|
||||||
|
model_field = Transaction._meta.get_field(field)
|
||||||
|
except FieldDoesNotExist:
|
||||||
|
return query.filter(**{f"{field}__in": clean_values})
|
||||||
|
|
||||||
|
if getattr(model_field, "many_to_many", False):
|
||||||
|
# For m2m fields (e.g., entities/tags), apply one filter per value so
|
||||||
|
# all provided values must be present in the matched transaction.
|
||||||
|
if all(self._is_int_like(v) for v in clean_values):
|
||||||
|
for value in clean_values:
|
||||||
|
query = query.filter(**{f"{field}__id": int(value)})
|
||||||
|
else:
|
||||||
|
for value in clean_values:
|
||||||
|
lookup = (
|
||||||
|
f"{field}__name"
|
||||||
|
if match_type == "strict"
|
||||||
|
else f"{field}__name__iexact"
|
||||||
|
)
|
||||||
|
query = query.filter(**{lookup: str(value).strip()})
|
||||||
|
|
||||||
|
return query.distinct()
|
||||||
|
|
||||||
|
return query.filter(**{f"{field}__in": clean_values})
|
||||||
|
|
||||||
def _coerce_type(
|
def _coerce_type(
|
||||||
self, value: str, mapping: version_1.ColumnMapping
|
self, value: str, mapping: version_1.ColumnMapping
|
||||||
) -> Union[str, int, bool, Decimal, datetime, list, None]:
|
) -> Union[str, int, bool, Decimal, datetime, list, None]:
|
||||||
@@ -844,6 +913,219 @@ class ImportService:
|
|||||||
f"Invalid {self.settings.file_type.upper()} file format: {str(e)}"
|
f"Invalid {self.settings.file_type.upper()} file format: {str(e)}"
|
||||||
)
|
)
|
||||||
|
|
||||||
|
def _parse_and_import_qif(self, content_lines: list[str], filename: str) -> None:
|
||||||
|
# Infer account from filename (remove extension)
|
||||||
|
account_name = os.path.splitext(os.path.basename(filename))[0]
|
||||||
|
|
||||||
|
current_transaction = {}
|
||||||
|
raw_lines_buffer = []
|
||||||
|
|
||||||
|
account = Account.objects.filter(name=account_name).first()
|
||||||
|
if not account:
|
||||||
|
raise ValueError(f"Account '{account_name}' not found.")
|
||||||
|
|
||||||
|
row_number = 0
|
||||||
|
for line in content_lines:
|
||||||
|
row_number += 1
|
||||||
|
line = line.strip()
|
||||||
|
if not line:
|
||||||
|
continue
|
||||||
|
|
||||||
|
raw_lines_buffer.append(line)
|
||||||
|
|
||||||
|
if line == "^":
|
||||||
|
if current_transaction:
|
||||||
|
# Deduplication using hash of raw lines
|
||||||
|
raw_content = "".join(raw_lines_buffer)
|
||||||
|
internal_id = hashlib.sha256(
|
||||||
|
raw_content.encode("utf-8")
|
||||||
|
).hexdigest()
|
||||||
|
|
||||||
|
# Reset buffer for next transaction
|
||||||
|
raw_lines_buffer = []
|
||||||
|
|
||||||
|
try:
|
||||||
|
with transaction.atomic():
|
||||||
|
if Transaction.objects.filter(
|
||||||
|
internal_id=internal_id
|
||||||
|
).exists():
|
||||||
|
self._increment_totals("skipped", 1)
|
||||||
|
self._log(
|
||||||
|
"info",
|
||||||
|
f"Skipped duplicate transaction from {filename}",
|
||||||
|
)
|
||||||
|
current_transaction = {}
|
||||||
|
continue
|
||||||
|
|
||||||
|
# Handle Account
|
||||||
|
if account:
|
||||||
|
current_transaction["account"] = account
|
||||||
|
else:
|
||||||
|
acc = Account.objects.filter(name=account_name).first()
|
||||||
|
if acc:
|
||||||
|
current_transaction["account"] = acc
|
||||||
|
else:
|
||||||
|
raise ValueError(
|
||||||
|
f"Account '{account_name}' not found."
|
||||||
|
)
|
||||||
|
|
||||||
|
current_transaction["internal_id"] = internal_id
|
||||||
|
|
||||||
|
# Handle Description/Memo mapping
|
||||||
|
if "memo" in current_transaction:
|
||||||
|
current_transaction["description"] = (
|
||||||
|
current_transaction.pop("memo")
|
||||||
|
)
|
||||||
|
|
||||||
|
# Handle Payee mapping
|
||||||
|
entities = []
|
||||||
|
if "payee" in current_transaction:
|
||||||
|
payee_name = current_transaction.pop("payee")
|
||||||
|
# "Treat the payee (P) as the entity. Use existing or create"
|
||||||
|
entity, _ = TransactionEntity.objects.get_or_create(
|
||||||
|
name=payee_name
|
||||||
|
)
|
||||||
|
entities.append(entity)
|
||||||
|
|
||||||
|
# Handle Label/Category
|
||||||
|
category = None
|
||||||
|
tags = []
|
||||||
|
if "label" in current_transaction:
|
||||||
|
label = current_transaction.pop("label")
|
||||||
|
if label.startswith("[") and label.endswith("]"):
|
||||||
|
# Transfer: set label as description, ignore category/tags
|
||||||
|
clean_label = label[1:-1]
|
||||||
|
current_transaction["description"] = clean_label
|
||||||
|
else:
|
||||||
|
parts = label.split(":")
|
||||||
|
if parts:
|
||||||
|
cat_name = parts[0].strip()
|
||||||
|
if cat_name:
|
||||||
|
category, _ = (
|
||||||
|
TransactionCategory.objects.get_or_create(
|
||||||
|
name=cat_name
|
||||||
|
)
|
||||||
|
)
|
||||||
|
|
||||||
|
if len(parts) > 1:
|
||||||
|
for tag_name in parts[1:]:
|
||||||
|
tag_name = tag_name.strip()
|
||||||
|
if tag_name:
|
||||||
|
tag, _ = (
|
||||||
|
TransactionTag.objects.get_or_create(
|
||||||
|
name=tag_name
|
||||||
|
)
|
||||||
|
)
|
||||||
|
tags.append(tag)
|
||||||
|
|
||||||
|
current_transaction["category"] = category
|
||||||
|
|
||||||
|
# Create transaction
|
||||||
|
new_trans = Transaction.objects.create(
|
||||||
|
**current_transaction
|
||||||
|
)
|
||||||
|
if entities:
|
||||||
|
new_trans.entities.set(entities)
|
||||||
|
if tags:
|
||||||
|
new_trans.tags.set(tags)
|
||||||
|
|
||||||
|
self.import_run.transactions.add(new_trans)
|
||||||
|
self._increment_totals("successful", 1)
|
||||||
|
|
||||||
|
except Exception as e:
|
||||||
|
if not self.settings.skip_errors:
|
||||||
|
raise e
|
||||||
|
self._log(
|
||||||
|
"warning",
|
||||||
|
f"Error processing transaction in {filename}: {str(e)}",
|
||||||
|
)
|
||||||
|
self._increment_totals("failed", 1)
|
||||||
|
|
||||||
|
# Reset for next transaction
|
||||||
|
current_transaction = {}
|
||||||
|
else:
|
||||||
|
# Empty transaction record (orphaned ^)
|
||||||
|
raw_lines_buffer = []
|
||||||
|
pass
|
||||||
|
self._increment_totals("processed", 1)
|
||||||
|
continue
|
||||||
|
|
||||||
|
if line.startswith("!"):
|
||||||
|
continue
|
||||||
|
|
||||||
|
code = line[0]
|
||||||
|
value = line[1:]
|
||||||
|
|
||||||
|
if code == "D":
|
||||||
|
try:
|
||||||
|
current_transaction["date"] = datetime.strptime(
|
||||||
|
value, self.settings.date_format
|
||||||
|
).date()
|
||||||
|
except ValueError:
|
||||||
|
self._log(
|
||||||
|
"warning",
|
||||||
|
f"Could not parse date '{value}' using format '{self.settings.date_format}' in {filename}",
|
||||||
|
)
|
||||||
|
if not self.settings.skip_errors:
|
||||||
|
raise ValueError(f"Invalid date format '{value}'")
|
||||||
|
|
||||||
|
elif code == "T":
|
||||||
|
try:
|
||||||
|
cleaned_value = value.replace(",", "")
|
||||||
|
amount = Decimal(cleaned_value)
|
||||||
|
if amount < 0:
|
||||||
|
current_transaction["type"] = Transaction.Type.EXPENSE
|
||||||
|
current_transaction["amount"] = abs(amount)
|
||||||
|
else:
|
||||||
|
current_transaction["type"] = Transaction.Type.INCOME
|
||||||
|
current_transaction["amount"] = amount
|
||||||
|
except InvalidOperation:
|
||||||
|
self._log(
|
||||||
|
"warning", f"Could not parse amount '{value}' in {filename}"
|
||||||
|
)
|
||||||
|
if not self.settings.skip_errors:
|
||||||
|
raise ValueError(f"Invalid amount format '{value}'")
|
||||||
|
|
||||||
|
elif code == "P":
|
||||||
|
current_transaction["payee"] = value
|
||||||
|
elif code == "M":
|
||||||
|
current_transaction["memo"] = value
|
||||||
|
elif code == "L":
|
||||||
|
current_transaction["label"] = value
|
||||||
|
elif code == "N":
|
||||||
|
pass
|
||||||
|
|
||||||
|
def _process_qif(self, file_path):
|
||||||
|
def process_logic():
|
||||||
|
if zipfile.is_zipfile(file_path):
|
||||||
|
try:
|
||||||
|
with zipfile.ZipFile(file_path, "r") as zf:
|
||||||
|
for filename in zf.namelist():
|
||||||
|
if filename.lower().endswith(
|
||||||
|
".qif"
|
||||||
|
) and not filename.startswith("__MACOSX"):
|
||||||
|
self._log(
|
||||||
|
"info", f"Processing QIF from ZIP: {filename}"
|
||||||
|
)
|
||||||
|
with zf.open(filename) as f:
|
||||||
|
content = f.read().decode(self.settings.encoding)
|
||||||
|
self._parse_and_import_qif(
|
||||||
|
content.splitlines(), filename
|
||||||
|
)
|
||||||
|
except Exception as e:
|
||||||
|
raise ValueError(f"Error processing ZIP file: {str(e)}")
|
||||||
|
else:
|
||||||
|
with open(file_path, "r", encoding=self.settings.encoding) as f:
|
||||||
|
self._parse_and_import_qif(
|
||||||
|
f.readlines(), os.path.basename(file_path)
|
||||||
|
)
|
||||||
|
|
||||||
|
if not self.settings.skip_errors:
|
||||||
|
with transaction.atomic():
|
||||||
|
process_logic()
|
||||||
|
else:
|
||||||
|
process_logic()
|
||||||
|
|
||||||
def _validate_file_path(self, file_path: str) -> str:
|
def _validate_file_path(self, file_path: str) -> str:
|
||||||
"""
|
"""
|
||||||
Validates that the file path is within the allowed temporary directory.
|
Validates that the file path is within the allowed temporary directory.
|
||||||
@@ -870,6 +1152,8 @@ class ImportService:
|
|||||||
self._process_csv(file_path)
|
self._process_csv(file_path)
|
||||||
elif isinstance(self.settings, version_1.ExcelImportSettings):
|
elif isinstance(self.settings, version_1.ExcelImportSettings):
|
||||||
self._process_excel(file_path)
|
self._process_excel(file_path)
|
||||||
|
elif isinstance(self.settings, version_1.QIFImportSettings):
|
||||||
|
self._process_qif(file_path)
|
||||||
|
|
||||||
self._update_status("FINISHED")
|
self._update_status("FINISHED")
|
||||||
self._log(
|
self._log(
|
||||||
|
|||||||
@@ -1,3 +0,0 @@
|
|||||||
from django.test import TestCase
|
|
||||||
|
|
||||||
# Create your tests here.
|
|
||||||
@@ -0,0 +1,311 @@
|
|||||||
|
"""
|
||||||
|
Tests for ImportService v1, specifically for deduplication logic.
|
||||||
|
|
||||||
|
These tests verify that the _check_duplicate_transaction method handles
|
||||||
|
different field types correctly, particularly ensuring that __iexact
|
||||||
|
is only used for string fields (not dates, decimals, etc.).
|
||||||
|
"""
|
||||||
|
|
||||||
|
from datetime import date
|
||||||
|
from decimal import Decimal
|
||||||
|
|
||||||
|
from django.test import TestCase
|
||||||
|
|
||||||
|
from apps.accounts.models import Account, AccountGroup
|
||||||
|
from apps.currencies.models import Currency
|
||||||
|
from apps.import_app.models import ImportProfile, ImportRun
|
||||||
|
from apps.import_app.services.v1 import ImportService
|
||||||
|
from apps.transactions.models import Transaction, TransactionEntity
|
||||||
|
|
||||||
|
|
||||||
|
class DeduplicationTests(TestCase):
|
||||||
|
"""Tests for transaction deduplication during import."""
|
||||||
|
|
||||||
|
def setUp(self):
|
||||||
|
"""Set up test data."""
|
||||||
|
self.currency = Currency.objects.create(
|
||||||
|
code="USD", name="US Dollar", decimal_places=2, prefix="$ "
|
||||||
|
)
|
||||||
|
self.account_group = AccountGroup.objects.create(name="Test Group")
|
||||||
|
self.account = Account.objects.create(
|
||||||
|
name="Test Account", group=self.account_group, currency=self.currency
|
||||||
|
)
|
||||||
|
|
||||||
|
# Create an existing transaction for deduplication tests
|
||||||
|
self.existing_transaction = Transaction.objects.create(
|
||||||
|
account=self.account,
|
||||||
|
type=Transaction.Type.EXPENSE,
|
||||||
|
date=date(2024, 1, 15),
|
||||||
|
amount=Decimal("100.00"),
|
||||||
|
description="Existing Transaction",
|
||||||
|
internal_id="ABC123",
|
||||||
|
)
|
||||||
|
|
||||||
|
def _create_import_service_with_deduplication(
|
||||||
|
self, fields: list[str], match_type: str = "lax"
|
||||||
|
) -> ImportService:
|
||||||
|
"""Helper to create an ImportService with specific deduplication rules."""
|
||||||
|
yaml_config = f"""
|
||||||
|
settings:
|
||||||
|
file_type: csv
|
||||||
|
importing: transactions
|
||||||
|
trigger_transaction_rules: false
|
||||||
|
mapping:
|
||||||
|
date_field:
|
||||||
|
source: date
|
||||||
|
target: date
|
||||||
|
format: "%Y-%m-%d"
|
||||||
|
amount_field:
|
||||||
|
source: amount
|
||||||
|
target: amount
|
||||||
|
description_field:
|
||||||
|
source: description
|
||||||
|
target: description
|
||||||
|
account_field:
|
||||||
|
source: account
|
||||||
|
target: account
|
||||||
|
type: id
|
||||||
|
deduplication:
|
||||||
|
- type: compare
|
||||||
|
fields: {fields}
|
||||||
|
match_type: {match_type}
|
||||||
|
"""
|
||||||
|
profile = ImportProfile.objects.create(
|
||||||
|
name=f"Test Profile {match_type} {'_'.join(fields)}",
|
||||||
|
yaml_config=yaml_config,
|
||||||
|
version=ImportProfile.Versions.VERSION_1,
|
||||||
|
)
|
||||||
|
import_run = ImportRun.objects.create(
|
||||||
|
profile=profile,
|
||||||
|
file_name="test.csv",
|
||||||
|
)
|
||||||
|
return ImportService(import_run)
|
||||||
|
|
||||||
|
def test_deduplication_with_date_field_strict_match(self):
|
||||||
|
"""Test that date fields work with strict matching."""
|
||||||
|
service = self._create_import_service_with_deduplication(
|
||||||
|
fields=["date"], match_type="strict"
|
||||||
|
)
|
||||||
|
|
||||||
|
# Should find duplicate when date matches
|
||||||
|
is_duplicate = service._check_duplicate_transaction({"date": date(2024, 1, 15)})
|
||||||
|
self.assertTrue(is_duplicate)
|
||||||
|
|
||||||
|
# Should not find duplicate when date differs
|
||||||
|
is_duplicate = service._check_duplicate_transaction({"date": date(2024, 2, 20)})
|
||||||
|
self.assertFalse(is_duplicate)
|
||||||
|
|
||||||
|
def test_deduplication_with_date_field_lax_match(self):
|
||||||
|
"""
|
||||||
|
Test that date fields use strict matching even when match_type is 'lax'.
|
||||||
|
|
||||||
|
This is the fix for the UPPER(date) PostgreSQL error. Date fields
|
||||||
|
cannot use __iexact, so they should fall back to strict matching.
|
||||||
|
"""
|
||||||
|
service = self._create_import_service_with_deduplication(
|
||||||
|
fields=["date"], match_type="lax"
|
||||||
|
)
|
||||||
|
|
||||||
|
# Should find duplicate when date matches (using strict comparison)
|
||||||
|
is_duplicate = service._check_duplicate_transaction({"date": date(2024, 1, 15)})
|
||||||
|
self.assertTrue(is_duplicate)
|
||||||
|
|
||||||
|
# Should not find duplicate when date differs
|
||||||
|
is_duplicate = service._check_duplicate_transaction({"date": date(2024, 2, 20)})
|
||||||
|
self.assertFalse(is_duplicate)
|
||||||
|
|
||||||
|
def test_deduplication_with_amount_field_lax_match(self):
|
||||||
|
"""
|
||||||
|
Test that Decimal fields use strict matching even when match_type is 'lax'.
|
||||||
|
|
||||||
|
Decimal fields cannot use __iexact, so they should fall back to strict matching.
|
||||||
|
"""
|
||||||
|
service = self._create_import_service_with_deduplication(
|
||||||
|
fields=["amount"], match_type="lax"
|
||||||
|
)
|
||||||
|
|
||||||
|
# Should find duplicate when amount matches
|
||||||
|
is_duplicate = service._check_duplicate_transaction(
|
||||||
|
{"amount": Decimal("100.00")}
|
||||||
|
)
|
||||||
|
self.assertTrue(is_duplicate)
|
||||||
|
|
||||||
|
# Should not find duplicate when amount differs
|
||||||
|
is_duplicate = service._check_duplicate_transaction(
|
||||||
|
{"amount": Decimal("200.00")}
|
||||||
|
)
|
||||||
|
self.assertFalse(is_duplicate)
|
||||||
|
|
||||||
|
def test_deduplication_with_string_field_lax_match(self):
|
||||||
|
"""
|
||||||
|
Test that string fields use case-insensitive matching with match_type 'lax'.
|
||||||
|
"""
|
||||||
|
service = self._create_import_service_with_deduplication(
|
||||||
|
fields=["description"], match_type="lax"
|
||||||
|
)
|
||||||
|
|
||||||
|
# Should find duplicate with case-insensitive match
|
||||||
|
is_duplicate = service._check_duplicate_transaction(
|
||||||
|
{"description": "EXISTING TRANSACTION"}
|
||||||
|
)
|
||||||
|
self.assertTrue(is_duplicate)
|
||||||
|
|
||||||
|
# Should find duplicate with exact case match
|
||||||
|
is_duplicate = service._check_duplicate_transaction(
|
||||||
|
{"description": "Existing Transaction"}
|
||||||
|
)
|
||||||
|
self.assertTrue(is_duplicate)
|
||||||
|
|
||||||
|
# Should not find duplicate when description differs
|
||||||
|
is_duplicate = service._check_duplicate_transaction(
|
||||||
|
{"description": "Different Transaction"}
|
||||||
|
)
|
||||||
|
self.assertFalse(is_duplicate)
|
||||||
|
|
||||||
|
def test_deduplication_with_string_field_strict_match(self):
|
||||||
|
"""
|
||||||
|
Test that string fields use case-sensitive matching with match_type 'strict'.
|
||||||
|
"""
|
||||||
|
service = self._create_import_service_with_deduplication(
|
||||||
|
fields=["description"], match_type="strict"
|
||||||
|
)
|
||||||
|
|
||||||
|
# Should NOT find duplicate with different case (strict matching)
|
||||||
|
is_duplicate = service._check_duplicate_transaction(
|
||||||
|
{"description": "EXISTING TRANSACTION"}
|
||||||
|
)
|
||||||
|
self.assertFalse(is_duplicate)
|
||||||
|
|
||||||
|
# Should find duplicate with exact case match
|
||||||
|
is_duplicate = service._check_duplicate_transaction(
|
||||||
|
{"description": "Existing Transaction"}
|
||||||
|
)
|
||||||
|
self.assertTrue(is_duplicate)
|
||||||
|
|
||||||
|
def test_deduplication_with_multiple_fields_mixed_types(self):
|
||||||
|
"""
|
||||||
|
Test deduplication with multiple fields of different types.
|
||||||
|
|
||||||
|
Verifies that string fields use __iexact while non-string fields
|
||||||
|
use strict matching, all in the same deduplication rule.
|
||||||
|
"""
|
||||||
|
service = self._create_import_service_with_deduplication(
|
||||||
|
fields=["date", "amount", "description"], match_type="lax"
|
||||||
|
)
|
||||||
|
|
||||||
|
# Should find duplicate when all fields match (with case-insensitive description)
|
||||||
|
is_duplicate = service._check_duplicate_transaction(
|
||||||
|
{
|
||||||
|
"date": date(2024, 1, 15),
|
||||||
|
"amount": Decimal("100.00"),
|
||||||
|
"description": "existing transaction", # lowercase should match
|
||||||
|
}
|
||||||
|
)
|
||||||
|
self.assertTrue(is_duplicate)
|
||||||
|
|
||||||
|
# Should NOT find duplicate when date differs
|
||||||
|
is_duplicate = service._check_duplicate_transaction(
|
||||||
|
{
|
||||||
|
"date": date(2024, 2, 20),
|
||||||
|
"amount": Decimal("100.00"),
|
||||||
|
"description": "existing transaction",
|
||||||
|
}
|
||||||
|
)
|
||||||
|
self.assertFalse(is_duplicate)
|
||||||
|
|
||||||
|
# Should NOT find duplicate when amount differs
|
||||||
|
is_duplicate = service._check_duplicate_transaction(
|
||||||
|
{
|
||||||
|
"date": date(2024, 1, 15),
|
||||||
|
"amount": Decimal("999.99"),
|
||||||
|
"description": "existing transaction",
|
||||||
|
}
|
||||||
|
)
|
||||||
|
self.assertFalse(is_duplicate)
|
||||||
|
|
||||||
|
def test_deduplication_with_internal_id_lax_match(self):
|
||||||
|
"""Test deduplication with internal_id field using lax matching."""
|
||||||
|
service = self._create_import_service_with_deduplication(
|
||||||
|
fields=["internal_id"], match_type="lax"
|
||||||
|
)
|
||||||
|
|
||||||
|
# Should find duplicate with case-insensitive match
|
||||||
|
is_duplicate = service._check_duplicate_transaction(
|
||||||
|
{"internal_id": "abc123"} # lowercase should match ABC123
|
||||||
|
)
|
||||||
|
self.assertTrue(is_duplicate)
|
||||||
|
|
||||||
|
# Should find duplicate with exact match
|
||||||
|
is_duplicate = service._check_duplicate_transaction({"internal_id": "ABC123"})
|
||||||
|
self.assertTrue(is_duplicate)
|
||||||
|
|
||||||
|
# Should not find duplicate when internal_id differs
|
||||||
|
is_duplicate = service._check_duplicate_transaction({"internal_id": "XYZ789"})
|
||||||
|
self.assertFalse(is_duplicate)
|
||||||
|
|
||||||
|
def test_no_duplicate_when_no_transactions_exist(self):
|
||||||
|
"""Test that no duplicate is found when there are no matching transactions."""
|
||||||
|
# Hard delete to bypass signals that require user context
|
||||||
|
self.existing_transaction.hard_delete()
|
||||||
|
|
||||||
|
service = self._create_import_service_with_deduplication(
|
||||||
|
fields=["date", "amount"], match_type="lax"
|
||||||
|
)
|
||||||
|
|
||||||
|
is_duplicate = service._check_duplicate_transaction(
|
||||||
|
{
|
||||||
|
"date": date(2024, 1, 15),
|
||||||
|
"amount": Decimal("100.00"),
|
||||||
|
}
|
||||||
|
)
|
||||||
|
self.assertFalse(is_duplicate)
|
||||||
|
|
||||||
|
def test_deduplication_with_missing_field_in_data(self):
|
||||||
|
"""Test that missing fields in transaction_data are handled gracefully."""
|
||||||
|
service = self._create_import_service_with_deduplication(
|
||||||
|
fields=["date", "nonexistent_field"], match_type="lax"
|
||||||
|
)
|
||||||
|
|
||||||
|
# Should still work, only checking the fields that exist
|
||||||
|
is_duplicate = service._check_duplicate_transaction(
|
||||||
|
{
|
||||||
|
"date": date(2024, 1, 15),
|
||||||
|
}
|
||||||
|
)
|
||||||
|
self.assertTrue(is_duplicate)
|
||||||
|
|
||||||
|
def test_deduplication_with_entities_list_value(self):
|
||||||
|
"""Test that list values for m2m entities deduplicate correctly."""
|
||||||
|
entity = TransactionEntity.objects.create(name="DB Vertrieb GmbH")
|
||||||
|
self.existing_transaction.entities.add(entity)
|
||||||
|
|
||||||
|
service = self._create_import_service_with_deduplication(
|
||||||
|
fields=["date", "amount", "entities"], match_type="strict"
|
||||||
|
)
|
||||||
|
|
||||||
|
is_duplicate = service._check_duplicate_transaction(
|
||||||
|
{
|
||||||
|
"date": date(2024, 1, 15),
|
||||||
|
"amount": Decimal("100.00"),
|
||||||
|
"entities": ["DB Vertrieb GmbH"],
|
||||||
|
}
|
||||||
|
)
|
||||||
|
self.assertTrue(is_duplicate)
|
||||||
|
|
||||||
|
def test_deduplication_with_entities_list_value_not_matching(self):
|
||||||
|
"""Test that non-matching entity list values are not marked duplicate."""
|
||||||
|
entity = TransactionEntity.objects.create(name="DB Vertrieb GmbH")
|
||||||
|
self.existing_transaction.entities.add(entity)
|
||||||
|
|
||||||
|
service = self._create_import_service_with_deduplication(
|
||||||
|
fields=["date", "amount", "entities"], match_type="strict"
|
||||||
|
)
|
||||||
|
|
||||||
|
is_duplicate = service._check_duplicate_transaction(
|
||||||
|
{
|
||||||
|
"date": date(2024, 1, 15),
|
||||||
|
"amount": Decimal("100.00"),
|
||||||
|
"entities": ["Different Entity"],
|
||||||
|
}
|
||||||
|
)
|
||||||
|
self.assertFalse(is_duplicate)
|
||||||
@@ -0,0 +1,259 @@
|
|||||||
|
from decimal import Decimal
|
||||||
|
import os
|
||||||
|
import shutil
|
||||||
|
from django.test import TestCase
|
||||||
|
from django.contrib.auth import get_user_model
|
||||||
|
from apps.accounts.models import Account, AccountGroup
|
||||||
|
from apps.currencies.models import Currency
|
||||||
|
from apps.common.middleware.thread_local import write_current_user, delete_current_user
|
||||||
|
from apps.import_app.models import ImportProfile, ImportRun
|
||||||
|
from apps.import_app.services.v1 import ImportService
|
||||||
|
from apps.transactions.models import (
|
||||||
|
Transaction,
|
||||||
|
)
|
||||||
|
|
||||||
|
|
||||||
|
class QIFImportTests(TestCase):
|
||||||
|
def setUp(self):
|
||||||
|
# Patch TEMP_DIR for testing
|
||||||
|
self.original_temp_dir = ImportService.TEMP_DIR
|
||||||
|
self.test_dir = os.path.abspath("temp_test_import")
|
||||||
|
ImportService.TEMP_DIR = self.test_dir
|
||||||
|
os.makedirs(self.test_dir, exist_ok=True)
|
||||||
|
|
||||||
|
# Create user and set context
|
||||||
|
User = get_user_model()
|
||||||
|
self.user = User.objects.create_user(
|
||||||
|
email="test@example.com", password="password"
|
||||||
|
)
|
||||||
|
write_current_user(self.user)
|
||||||
|
|
||||||
|
self.currency = Currency.objects.create(
|
||||||
|
code="BRL", name="Real", decimal_places=2, prefix="R$ "
|
||||||
|
)
|
||||||
|
self.group = AccountGroup.objects.create(name="Test Group", owner=self.user)
|
||||||
|
self.account = Account.objects.create(
|
||||||
|
name="bradesco-checking",
|
||||||
|
group=self.group,
|
||||||
|
currency=self.currency,
|
||||||
|
owner=self.user,
|
||||||
|
)
|
||||||
|
|
||||||
|
def tearDown(self):
|
||||||
|
delete_current_user()
|
||||||
|
ImportService.TEMP_DIR = self.original_temp_dir
|
||||||
|
if os.path.exists(self.test_dir):
|
||||||
|
shutil.rmtree(self.test_dir)
|
||||||
|
|
||||||
|
def test_import_single_qif_valid_mapping(self):
|
||||||
|
content = """!Type:Bank
|
||||||
|
D04/01/2015
|
||||||
|
T8069.46
|
||||||
|
PMy Payee -> Entity
|
||||||
|
MNote -> Desc
|
||||||
|
LOld Cat:New Tag
|
||||||
|
^
|
||||||
|
D05/01/2015
|
||||||
|
T-100.00
|
||||||
|
PSupermarket
|
||||||
|
MWeekly shopping
|
||||||
|
L[Transfer]
|
||||||
|
^
|
||||||
|
"""
|
||||||
|
filename = "bradesco-checking.qif"
|
||||||
|
file_path = os.path.join(self.test_dir, filename)
|
||||||
|
with open(file_path, "w", encoding="utf-8") as f:
|
||||||
|
f.write(content)
|
||||||
|
|
||||||
|
yaml_config = """
|
||||||
|
settings:
|
||||||
|
file_type: qif
|
||||||
|
importing: transactions
|
||||||
|
date_format: "%d/%m/%Y"
|
||||||
|
mapping: {}
|
||||||
|
"""
|
||||||
|
profile = ImportProfile.objects.create(
|
||||||
|
name="QIF Profile",
|
||||||
|
yaml_config=yaml_config,
|
||||||
|
version=ImportProfile.Versions.VERSION_1,
|
||||||
|
)
|
||||||
|
run = ImportRun.objects.create(profile=profile, file_name=filename)
|
||||||
|
service = ImportService(run)
|
||||||
|
|
||||||
|
service.process_file(file_path)
|
||||||
|
|
||||||
|
self.assertEqual(Transaction.objects.count(), 2)
|
||||||
|
|
||||||
|
# Transaction 1: Income, Category+Tag
|
||||||
|
t1 = Transaction.objects.get(description="Note -> Desc")
|
||||||
|
self.assertEqual(t1.amount, Decimal("8069.46"))
|
||||||
|
self.assertEqual(t1.type, Transaction.Type.INCOME)
|
||||||
|
self.assertEqual(t1.category.name, "Old Cat")
|
||||||
|
self.assertTrue(t1.tags.filter(name="New Tag").exists())
|
||||||
|
self.assertTrue(t1.entities.filter(name="My Payee -> Entity").exists())
|
||||||
|
self.assertEqual(t1.account, self.account)
|
||||||
|
|
||||||
|
# Transaction 2: Expense, Transfer ([Transfer] -> Description)
|
||||||
|
t2 = Transaction.objects.get(description="Transfer")
|
||||||
|
self.assertEqual(t2.amount, Decimal("100.00"))
|
||||||
|
self.assertEqual(t2.type, Transaction.Type.EXPENSE)
|
||||||
|
self.assertIsNone(t2.category)
|
||||||
|
self.assertFalse(t2.tags.exists())
|
||||||
|
self.assertTrue(t2.entities.filter(name="Supermarket").exists())
|
||||||
|
self.assertEqual(t2.description, "Transfer")
|
||||||
|
|
||||||
|
def test_import_deduplication_hash(self):
|
||||||
|
# Same content twice. Should result in only 1 transaction due to hash deduplication.
|
||||||
|
content = """!Type:Bank
|
||||||
|
D04/01/2015
|
||||||
|
T100.00
|
||||||
|
POK
|
||||||
|
^
|
||||||
|
"""
|
||||||
|
filename = "bradesco-checking.qif"
|
||||||
|
file_path = os.path.join(self.test_dir, filename)
|
||||||
|
with open(file_path, "w", encoding="utf-8") as f:
|
||||||
|
f.write(content)
|
||||||
|
|
||||||
|
yaml_config = """
|
||||||
|
settings:
|
||||||
|
file_type: qif
|
||||||
|
importing: transactions
|
||||||
|
date_format: "%d/%m/%Y"
|
||||||
|
mapping: {}
|
||||||
|
"""
|
||||||
|
profile = ImportProfile.objects.create(
|
||||||
|
name="QIF Profile",
|
||||||
|
yaml_config=yaml_config,
|
||||||
|
version=ImportProfile.Versions.VERSION_1,
|
||||||
|
)
|
||||||
|
run = ImportRun.objects.create(profile=profile, file_name=filename)
|
||||||
|
service = ImportService(run)
|
||||||
|
|
||||||
|
# First run
|
||||||
|
service.process_file(file_path)
|
||||||
|
self.assertEqual(Transaction.objects.count(), 1)
|
||||||
|
|
||||||
|
# Service deletes file after processing, so recreate it for second run
|
||||||
|
with open(file_path, "w", encoding="utf-8") as f:
|
||||||
|
f.write(content)
|
||||||
|
|
||||||
|
# Second run - Duplicate content
|
||||||
|
service.process_file(file_path)
|
||||||
|
self.assertEqual(Transaction.objects.count(), 1)
|
||||||
|
|
||||||
|
def test_import_strict_error_rollback(self):
|
||||||
|
# atomic check.
|
||||||
|
# Transaction 1 valid, Transaction 2 invalid date.
|
||||||
|
content = """!Type:Bank
|
||||||
|
D04/01/2015
|
||||||
|
T100.00
|
||||||
|
POK
|
||||||
|
^
|
||||||
|
DINVALID
|
||||||
|
T100.00
|
||||||
|
PBad
|
||||||
|
^
|
||||||
|
"""
|
||||||
|
filename = "bradesco-checking.qif"
|
||||||
|
file_path = os.path.join(self.test_dir, filename)
|
||||||
|
with open(file_path, "w", encoding="utf-8") as f:
|
||||||
|
f.write(content)
|
||||||
|
|
||||||
|
yaml_config = """
|
||||||
|
settings:
|
||||||
|
file_type: qif
|
||||||
|
importing: transactions
|
||||||
|
date_format: "%d/%m/%Y"
|
||||||
|
skip_errors: false
|
||||||
|
mapping: {}
|
||||||
|
"""
|
||||||
|
profile = ImportProfile.objects.create(
|
||||||
|
name="QIF Profile",
|
||||||
|
yaml_config=yaml_config,
|
||||||
|
version=ImportProfile.Versions.VERSION_1,
|
||||||
|
)
|
||||||
|
run = ImportRun.objects.create(profile=profile, file_name=filename)
|
||||||
|
service = ImportService(run)
|
||||||
|
|
||||||
|
with self.assertRaises(Exception) as cm:
|
||||||
|
service.process_file(file_path)
|
||||||
|
self.assertEqual(str(cm.exception), "Import failed")
|
||||||
|
|
||||||
|
# Should be 0 transactions because of atomic rollback
|
||||||
|
self.assertEqual(Transaction.objects.count(), 0)
|
||||||
|
|
||||||
|
def test_import_missing_account(self):
|
||||||
|
# File with account name that doesn't exist
|
||||||
|
content = """!Type:Bank
|
||||||
|
D04/01/2015
|
||||||
|
T100.00
|
||||||
|
POK
|
||||||
|
^
|
||||||
|
"""
|
||||||
|
filename = "missing-account.qif"
|
||||||
|
file_path = os.path.join(self.test_dir, filename)
|
||||||
|
with open(file_path, "w", encoding="utf-8") as f:
|
||||||
|
f.write(content)
|
||||||
|
|
||||||
|
yaml_config = """
|
||||||
|
settings:
|
||||||
|
file_type: qif
|
||||||
|
importing: transactions
|
||||||
|
date_format: "%d/%m/%Y"
|
||||||
|
mapping: {}
|
||||||
|
"""
|
||||||
|
profile = ImportProfile.objects.create(
|
||||||
|
name="QIF Profile",
|
||||||
|
yaml_config=yaml_config,
|
||||||
|
version=ImportProfile.Versions.VERSION_1,
|
||||||
|
)
|
||||||
|
run = ImportRun.objects.create(profile=profile, file_name=filename)
|
||||||
|
service = ImportService(run)
|
||||||
|
|
||||||
|
# Should fail because account doesn't exist
|
||||||
|
with self.assertRaises(Exception) as cm:
|
||||||
|
service.process_file(file_path)
|
||||||
|
self.assertEqual(str(cm.exception), "Import failed")
|
||||||
|
|
||||||
|
def test_import_skip_errors(self):
|
||||||
|
# skip_errors: true.
|
||||||
|
# Transaction 1 valid, Transaction 2 invalid date.
|
||||||
|
content = """!Type:Bank
|
||||||
|
D04/01/2015
|
||||||
|
T100.00
|
||||||
|
POK
|
||||||
|
^
|
||||||
|
DINVALID
|
||||||
|
T100.00
|
||||||
|
PBad
|
||||||
|
^
|
||||||
|
"""
|
||||||
|
filename = "bradesco-checking.qif"
|
||||||
|
file_path = os.path.join(self.test_dir, filename)
|
||||||
|
with open(file_path, "w", encoding="utf-8") as f:
|
||||||
|
f.write(content)
|
||||||
|
|
||||||
|
yaml_config = """
|
||||||
|
settings:
|
||||||
|
file_type: qif
|
||||||
|
importing: transactions
|
||||||
|
date_format: "%d/%m/%Y"
|
||||||
|
skip_errors: true
|
||||||
|
mapping: {}
|
||||||
|
"""
|
||||||
|
profile = ImportProfile.objects.create(
|
||||||
|
name="QIF Profile",
|
||||||
|
yaml_config=yaml_config,
|
||||||
|
version=ImportProfile.Versions.VERSION_1,
|
||||||
|
)
|
||||||
|
run = ImportRun.objects.create(profile=profile, file_name=filename)
|
||||||
|
service = ImportService(run)
|
||||||
|
|
||||||
|
service.process_file(file_path)
|
||||||
|
|
||||||
|
# Should be 1 transaction (valid one)
|
||||||
|
self.assertEqual(Transaction.objects.count(), 1)
|
||||||
|
self.assertEqual(
|
||||||
|
Transaction.objects.first().description, ""
|
||||||
|
) # empty desc if no memo
|
||||||
@@ -49,4 +49,14 @@ urlpatterns = [
|
|||||||
views.emergency_fund,
|
views.emergency_fund,
|
||||||
name="insights_emergency_fund",
|
name="insights_emergency_fund",
|
||||||
),
|
),
|
||||||
|
path(
|
||||||
|
"insights/year-by-year/",
|
||||||
|
views.year_by_year,
|
||||||
|
name="insights_year_by_year",
|
||||||
|
),
|
||||||
|
path(
|
||||||
|
"insights/month-by-month/",
|
||||||
|
views.month_by_month,
|
||||||
|
name="insights_month_by_month",
|
||||||
|
),
|
||||||
]
|
]
|
||||||
|
|||||||
@@ -0,0 +1,316 @@
|
|||||||
|
from collections import OrderedDict
|
||||||
|
from decimal import Decimal
|
||||||
|
|
||||||
|
from django.db import models
|
||||||
|
from django.db.models import Sum, Case, When, Value
|
||||||
|
from django.db.models.functions import Coalesce
|
||||||
|
from django.utils import timezone
|
||||||
|
|
||||||
|
from apps.currencies.models import Currency
|
||||||
|
from apps.currencies.utils.convert import convert
|
||||||
|
from apps.transactions.models import Transaction
|
||||||
|
|
||||||
|
|
||||||
|
def get_month_by_month_data(year=None, group_by="categories"):
|
||||||
|
"""
|
||||||
|
Aggregate transaction totals by month for a specific year, grouped by categories, tags, or entities.
|
||||||
|
|
||||||
|
Args:
|
||||||
|
year: The year to filter transactions (defaults to current year)
|
||||||
|
group_by: One of "categories", "tags", or "entities"
|
||||||
|
|
||||||
|
Returns:
|
||||||
|
{
|
||||||
|
"year": 2025,
|
||||||
|
"available_years": [2025, 2024, ...],
|
||||||
|
"months": [1, 2, 3, ..., 12],
|
||||||
|
"items": {
|
||||||
|
item_id: {
|
||||||
|
"name": "Item Name",
|
||||||
|
"month_totals": {
|
||||||
|
1: {"currencies": {...}},
|
||||||
|
...
|
||||||
|
},
|
||||||
|
"total": {"currencies": {...}}
|
||||||
|
},
|
||||||
|
...
|
||||||
|
},
|
||||||
|
"month_totals": {...},
|
||||||
|
"grand_total": {"currencies": {...}}
|
||||||
|
}
|
||||||
|
"""
|
||||||
|
if year is None:
|
||||||
|
year = timezone.localdate(timezone.now()).year
|
||||||
|
|
||||||
|
# Base queryset - all paid transactions, non-muted
|
||||||
|
transactions = Transaction.objects.filter(
|
||||||
|
is_paid=True,
|
||||||
|
account__is_archived=False,
|
||||||
|
).exclude(account__currency__is_archived=True)
|
||||||
|
|
||||||
|
# Get available years for the selector
|
||||||
|
available_years = list(
|
||||||
|
transactions.values_list("reference_date__year", flat=True)
|
||||||
|
.distinct()
|
||||||
|
.order_by("-reference_date__year")
|
||||||
|
)
|
||||||
|
|
||||||
|
# Filter by the selected year
|
||||||
|
transactions = transactions.filter(reference_date__year=year)
|
||||||
|
|
||||||
|
# Define grouping fields based on group_by parameter
|
||||||
|
if group_by == "tags":
|
||||||
|
group_field = "tags"
|
||||||
|
name_field = "tags__name"
|
||||||
|
elif group_by == "entities":
|
||||||
|
group_field = "entities"
|
||||||
|
name_field = "entities__name"
|
||||||
|
else: # Default to categories
|
||||||
|
group_field = "category"
|
||||||
|
name_field = "category__name"
|
||||||
|
|
||||||
|
# Months 1-12
|
||||||
|
months = list(range(1, 13))
|
||||||
|
|
||||||
|
if not available_years:
|
||||||
|
return {
|
||||||
|
"year": year,
|
||||||
|
"available_years": [],
|
||||||
|
"months": months,
|
||||||
|
"items": {},
|
||||||
|
"month_totals": {},
|
||||||
|
"grand_total": {"currencies": {}},
|
||||||
|
}
|
||||||
|
|
||||||
|
# Aggregate by group, month, and currency
|
||||||
|
metrics = (
|
||||||
|
transactions.values(
|
||||||
|
group_field,
|
||||||
|
name_field,
|
||||||
|
"reference_date__month",
|
||||||
|
"account__currency",
|
||||||
|
"account__currency__code",
|
||||||
|
"account__currency__name",
|
||||||
|
"account__currency__decimal_places",
|
||||||
|
"account__currency__prefix",
|
||||||
|
"account__currency__suffix",
|
||||||
|
"account__currency__exchange_currency",
|
||||||
|
)
|
||||||
|
.annotate(
|
||||||
|
expense_total=Coalesce(
|
||||||
|
Sum(
|
||||||
|
Case(
|
||||||
|
When(type=Transaction.Type.EXPENSE, then="amount"),
|
||||||
|
default=Value(0),
|
||||||
|
output_field=models.DecimalField(),
|
||||||
|
)
|
||||||
|
),
|
||||||
|
Decimal("0"),
|
||||||
|
),
|
||||||
|
income_total=Coalesce(
|
||||||
|
Sum(
|
||||||
|
Case(
|
||||||
|
When(type=Transaction.Type.INCOME, then="amount"),
|
||||||
|
default=Value(0),
|
||||||
|
output_field=models.DecimalField(),
|
||||||
|
)
|
||||||
|
),
|
||||||
|
Decimal("0"),
|
||||||
|
),
|
||||||
|
)
|
||||||
|
.order_by(name_field, "reference_date__month")
|
||||||
|
)
|
||||||
|
|
||||||
|
# Build result structure
|
||||||
|
result = {
|
||||||
|
"year": year,
|
||||||
|
"available_years": available_years,
|
||||||
|
"months": months,
|
||||||
|
"items": OrderedDict(),
|
||||||
|
"month_totals": {},
|
||||||
|
"grand_total": {"currencies": {}},
|
||||||
|
}
|
||||||
|
|
||||||
|
# Store currency info for later use in totals
|
||||||
|
currency_info = {}
|
||||||
|
|
||||||
|
for metric in metrics:
|
||||||
|
item_id = metric[group_field]
|
||||||
|
item_name = metric[name_field]
|
||||||
|
month = metric["reference_date__month"]
|
||||||
|
currency_id = metric["account__currency"]
|
||||||
|
|
||||||
|
# Use a consistent key for None (uncategorized/untagged/no entity)
|
||||||
|
item_key = item_id if item_id is not None else "__none__"
|
||||||
|
|
||||||
|
if item_key not in result["items"]:
|
||||||
|
result["items"][item_key] = {
|
||||||
|
"name": item_name,
|
||||||
|
"month_totals": {},
|
||||||
|
"total": {"currencies": {}},
|
||||||
|
}
|
||||||
|
|
||||||
|
if month not in result["items"][item_key]["month_totals"]:
|
||||||
|
result["items"][item_key]["month_totals"][month] = {"currencies": {}}
|
||||||
|
|
||||||
|
# Calculate final total (income - expense)
|
||||||
|
final_total = metric["income_total"] - metric["expense_total"]
|
||||||
|
|
||||||
|
# Store currency info for totals calculation
|
||||||
|
if currency_id not in currency_info:
|
||||||
|
currency_info[currency_id] = {
|
||||||
|
"code": metric["account__currency__code"],
|
||||||
|
"name": metric["account__currency__name"],
|
||||||
|
"decimal_places": metric["account__currency__decimal_places"],
|
||||||
|
"prefix": metric["account__currency__prefix"],
|
||||||
|
"suffix": metric["account__currency__suffix"],
|
||||||
|
"exchange_currency_id": metric["account__currency__exchange_currency"],
|
||||||
|
}
|
||||||
|
|
||||||
|
currency_data = {
|
||||||
|
"currency": {
|
||||||
|
"code": metric["account__currency__code"],
|
||||||
|
"name": metric["account__currency__name"],
|
||||||
|
"decimal_places": metric["account__currency__decimal_places"],
|
||||||
|
"prefix": metric["account__currency__prefix"],
|
||||||
|
"suffix": metric["account__currency__suffix"],
|
||||||
|
},
|
||||||
|
"final_total": final_total,
|
||||||
|
"income_total": metric["income_total"],
|
||||||
|
"expense_total": metric["expense_total"],
|
||||||
|
}
|
||||||
|
|
||||||
|
# Handle currency conversion if exchange currency is set
|
||||||
|
if metric["account__currency__exchange_currency"]:
|
||||||
|
from_currency = Currency.objects.get(id=currency_id)
|
||||||
|
exchange_currency = Currency.objects.get(
|
||||||
|
id=metric["account__currency__exchange_currency"]
|
||||||
|
)
|
||||||
|
|
||||||
|
converted_amount, prefix, suffix, decimal_places = convert(
|
||||||
|
amount=final_total,
|
||||||
|
from_currency=from_currency,
|
||||||
|
to_currency=exchange_currency,
|
||||||
|
)
|
||||||
|
|
||||||
|
if converted_amount is not None:
|
||||||
|
currency_data["exchanged"] = {
|
||||||
|
"final_total": converted_amount,
|
||||||
|
"currency": {
|
||||||
|
"prefix": prefix,
|
||||||
|
"suffix": suffix,
|
||||||
|
"decimal_places": decimal_places,
|
||||||
|
"code": exchange_currency.code,
|
||||||
|
"name": exchange_currency.name,
|
||||||
|
},
|
||||||
|
}
|
||||||
|
|
||||||
|
result["items"][item_key]["month_totals"][month]["currencies"][currency_id] = (
|
||||||
|
currency_data
|
||||||
|
)
|
||||||
|
|
||||||
|
# Accumulate item total (across all months for this item)
|
||||||
|
if currency_id not in result["items"][item_key]["total"]["currencies"]:
|
||||||
|
result["items"][item_key]["total"]["currencies"][currency_id] = {
|
||||||
|
"currency": currency_data["currency"].copy(),
|
||||||
|
"final_total": Decimal("0"),
|
||||||
|
}
|
||||||
|
result["items"][item_key]["total"]["currencies"][currency_id][
|
||||||
|
"final_total"
|
||||||
|
] += final_total
|
||||||
|
|
||||||
|
# Accumulate month total (across all items for this month)
|
||||||
|
if month not in result["month_totals"]:
|
||||||
|
result["month_totals"][month] = {"currencies": {}}
|
||||||
|
if currency_id not in result["month_totals"][month]["currencies"]:
|
||||||
|
result["month_totals"][month]["currencies"][currency_id] = {
|
||||||
|
"currency": currency_data["currency"].copy(),
|
||||||
|
"final_total": Decimal("0"),
|
||||||
|
}
|
||||||
|
result["month_totals"][month]["currencies"][currency_id]["final_total"] += (
|
||||||
|
final_total
|
||||||
|
)
|
||||||
|
|
||||||
|
# Accumulate grand total
|
||||||
|
if currency_id not in result["grand_total"]["currencies"]:
|
||||||
|
result["grand_total"]["currencies"][currency_id] = {
|
||||||
|
"currency": currency_data["currency"].copy(),
|
||||||
|
"final_total": Decimal("0"),
|
||||||
|
}
|
||||||
|
result["grand_total"]["currencies"][currency_id]["final_total"] += final_total
|
||||||
|
|
||||||
|
# Add currency conversion for item totals
|
||||||
|
for item_key, item_data in result["items"].items():
|
||||||
|
for currency_id, total_data in item_data["total"]["currencies"].items():
|
||||||
|
if currency_info[currency_id]["exchange_currency_id"]:
|
||||||
|
from_currency = Currency.objects.get(id=currency_id)
|
||||||
|
exchange_currency = Currency.objects.get(
|
||||||
|
id=currency_info[currency_id]["exchange_currency_id"]
|
||||||
|
)
|
||||||
|
converted_amount, prefix, suffix, decimal_places = convert(
|
||||||
|
amount=total_data["final_total"],
|
||||||
|
from_currency=from_currency,
|
||||||
|
to_currency=exchange_currency,
|
||||||
|
)
|
||||||
|
if converted_amount is not None:
|
||||||
|
total_data["exchanged"] = {
|
||||||
|
"final_total": converted_amount,
|
||||||
|
"currency": {
|
||||||
|
"prefix": prefix,
|
||||||
|
"suffix": suffix,
|
||||||
|
"decimal_places": decimal_places,
|
||||||
|
"code": exchange_currency.code,
|
||||||
|
"name": exchange_currency.name,
|
||||||
|
},
|
||||||
|
}
|
||||||
|
|
||||||
|
# Add currency conversion for month totals
|
||||||
|
for month, month_data in result["month_totals"].items():
|
||||||
|
for currency_id, total_data in month_data["currencies"].items():
|
||||||
|
if currency_info[currency_id]["exchange_currency_id"]:
|
||||||
|
from_currency = Currency.objects.get(id=currency_id)
|
||||||
|
exchange_currency = Currency.objects.get(
|
||||||
|
id=currency_info[currency_id]["exchange_currency_id"]
|
||||||
|
)
|
||||||
|
converted_amount, prefix, suffix, decimal_places = convert(
|
||||||
|
amount=total_data["final_total"],
|
||||||
|
from_currency=from_currency,
|
||||||
|
to_currency=exchange_currency,
|
||||||
|
)
|
||||||
|
if converted_amount is not None:
|
||||||
|
total_data["exchanged"] = {
|
||||||
|
"final_total": converted_amount,
|
||||||
|
"currency": {
|
||||||
|
"prefix": prefix,
|
||||||
|
"suffix": suffix,
|
||||||
|
"decimal_places": decimal_places,
|
||||||
|
"code": exchange_currency.code,
|
||||||
|
"name": exchange_currency.name,
|
||||||
|
},
|
||||||
|
}
|
||||||
|
|
||||||
|
# Add currency conversion for grand total
|
||||||
|
for currency_id, total_data in result["grand_total"]["currencies"].items():
|
||||||
|
if currency_info[currency_id]["exchange_currency_id"]:
|
||||||
|
from_currency = Currency.objects.get(id=currency_id)
|
||||||
|
exchange_currency = Currency.objects.get(
|
||||||
|
id=currency_info[currency_id]["exchange_currency_id"]
|
||||||
|
)
|
||||||
|
converted_amount, prefix, suffix, decimal_places = convert(
|
||||||
|
amount=total_data["final_total"],
|
||||||
|
from_currency=from_currency,
|
||||||
|
to_currency=exchange_currency,
|
||||||
|
)
|
||||||
|
if converted_amount is not None:
|
||||||
|
total_data["exchanged"] = {
|
||||||
|
"final_total": converted_amount,
|
||||||
|
"currency": {
|
||||||
|
"prefix": prefix,
|
||||||
|
"suffix": suffix,
|
||||||
|
"decimal_places": decimal_places,
|
||||||
|
"code": exchange_currency.code,
|
||||||
|
"name": exchange_currency.name,
|
||||||
|
},
|
||||||
|
}
|
||||||
|
|
||||||
|
return result
|
||||||
@@ -0,0 +1,303 @@
|
|||||||
|
from collections import OrderedDict
|
||||||
|
from decimal import Decimal
|
||||||
|
|
||||||
|
from django.db import models
|
||||||
|
from django.db.models import Sum, Case, When, Value
|
||||||
|
from django.db.models.functions import Coalesce
|
||||||
|
|
||||||
|
from apps.currencies.models import Currency
|
||||||
|
from apps.currencies.utils.convert import convert
|
||||||
|
from apps.transactions.models import Transaction
|
||||||
|
|
||||||
|
|
||||||
|
def get_year_by_year_data(group_by="categories"):
|
||||||
|
"""
|
||||||
|
Aggregate transaction totals by year for categories, tags, or entities.
|
||||||
|
|
||||||
|
Args:
|
||||||
|
group_by: One of "categories", "tags", or "entities"
|
||||||
|
|
||||||
|
Returns:
|
||||||
|
{
|
||||||
|
"years": [2025, 2024, ...], # Sorted descending
|
||||||
|
"items": {
|
||||||
|
item_id: {
|
||||||
|
"name": "Item Name",
|
||||||
|
"year_totals": {
|
||||||
|
2025: {"currencies": {...}},
|
||||||
|
...
|
||||||
|
},
|
||||||
|
"total": {"currencies": {...}} # Sum across all years
|
||||||
|
},
|
||||||
|
...
|
||||||
|
},
|
||||||
|
"year_totals": { # Sum across all items for each year
|
||||||
|
2025: {"currencies": {...}},
|
||||||
|
...
|
||||||
|
},
|
||||||
|
"grand_total": {"currencies": {...}} # Sum of everything
|
||||||
|
}
|
||||||
|
"""
|
||||||
|
# Base queryset - all paid transactions, non-muted
|
||||||
|
transactions = Transaction.objects.filter(
|
||||||
|
is_paid=True,
|
||||||
|
account__is_archived=False,
|
||||||
|
).exclude(account__currency__is_archived=True)
|
||||||
|
|
||||||
|
# Define grouping fields based on group_by parameter
|
||||||
|
if group_by == "tags":
|
||||||
|
group_field = "tags"
|
||||||
|
name_field = "tags__name"
|
||||||
|
elif group_by == "entities":
|
||||||
|
group_field = "entities"
|
||||||
|
name_field = "entities__name"
|
||||||
|
else: # Default to categories
|
||||||
|
group_field = "category"
|
||||||
|
name_field = "category__name"
|
||||||
|
|
||||||
|
# Get all unique years with transactions
|
||||||
|
years = (
|
||||||
|
transactions.values_list("reference_date__year", flat=True)
|
||||||
|
.distinct()
|
||||||
|
.order_by("-reference_date__year")
|
||||||
|
)
|
||||||
|
years = list(years)
|
||||||
|
|
||||||
|
if not years:
|
||||||
|
return {
|
||||||
|
"years": [],
|
||||||
|
"items": {},
|
||||||
|
"year_totals": {},
|
||||||
|
"grand_total": {"currencies": {}},
|
||||||
|
}
|
||||||
|
|
||||||
|
# Aggregate by group, year, and currency
|
||||||
|
metrics = (
|
||||||
|
transactions.values(
|
||||||
|
group_field,
|
||||||
|
name_field,
|
||||||
|
"reference_date__year",
|
||||||
|
"account__currency",
|
||||||
|
"account__currency__code",
|
||||||
|
"account__currency__name",
|
||||||
|
"account__currency__decimal_places",
|
||||||
|
"account__currency__prefix",
|
||||||
|
"account__currency__suffix",
|
||||||
|
"account__currency__exchange_currency",
|
||||||
|
)
|
||||||
|
.annotate(
|
||||||
|
expense_total=Coalesce(
|
||||||
|
Sum(
|
||||||
|
Case(
|
||||||
|
When(type=Transaction.Type.EXPENSE, then="amount"),
|
||||||
|
default=Value(0),
|
||||||
|
output_field=models.DecimalField(),
|
||||||
|
)
|
||||||
|
),
|
||||||
|
Decimal("0"),
|
||||||
|
),
|
||||||
|
income_total=Coalesce(
|
||||||
|
Sum(
|
||||||
|
Case(
|
||||||
|
When(type=Transaction.Type.INCOME, then="amount"),
|
||||||
|
default=Value(0),
|
||||||
|
output_field=models.DecimalField(),
|
||||||
|
)
|
||||||
|
),
|
||||||
|
Decimal("0"),
|
||||||
|
),
|
||||||
|
)
|
||||||
|
.order_by(name_field, "-reference_date__year")
|
||||||
|
)
|
||||||
|
|
||||||
|
# Build result structure
|
||||||
|
result = {
|
||||||
|
"years": years,
|
||||||
|
"items": OrderedDict(),
|
||||||
|
"year_totals": {}, # Totals per year across all items
|
||||||
|
"grand_total": {"currencies": {}}, # Grand total across everything
|
||||||
|
}
|
||||||
|
|
||||||
|
# Store currency info for later use in totals
|
||||||
|
currency_info = {}
|
||||||
|
|
||||||
|
for metric in metrics:
|
||||||
|
item_id = metric[group_field]
|
||||||
|
item_name = metric[name_field]
|
||||||
|
year = metric["reference_date__year"]
|
||||||
|
currency_id = metric["account__currency"]
|
||||||
|
|
||||||
|
# Use a consistent key for None (uncategorized/untagged/no entity)
|
||||||
|
item_key = item_id if item_id is not None else "__none__"
|
||||||
|
|
||||||
|
if item_key not in result["items"]:
|
||||||
|
result["items"][item_key] = {
|
||||||
|
"name": item_name,
|
||||||
|
"year_totals": {},
|
||||||
|
"total": {"currencies": {}}, # Total for this item across all years
|
||||||
|
}
|
||||||
|
|
||||||
|
if year not in result["items"][item_key]["year_totals"]:
|
||||||
|
result["items"][item_key]["year_totals"][year] = {"currencies": {}}
|
||||||
|
|
||||||
|
# Calculate final total (income - expense)
|
||||||
|
final_total = metric["income_total"] - metric["expense_total"]
|
||||||
|
|
||||||
|
# Store currency info for totals calculation
|
||||||
|
if currency_id not in currency_info:
|
||||||
|
currency_info[currency_id] = {
|
||||||
|
"code": metric["account__currency__code"],
|
||||||
|
"name": metric["account__currency__name"],
|
||||||
|
"decimal_places": metric["account__currency__decimal_places"],
|
||||||
|
"prefix": metric["account__currency__prefix"],
|
||||||
|
"suffix": metric["account__currency__suffix"],
|
||||||
|
"exchange_currency_id": metric["account__currency__exchange_currency"],
|
||||||
|
}
|
||||||
|
|
||||||
|
currency_data = {
|
||||||
|
"currency": {
|
||||||
|
"code": metric["account__currency__code"],
|
||||||
|
"name": metric["account__currency__name"],
|
||||||
|
"decimal_places": metric["account__currency__decimal_places"],
|
||||||
|
"prefix": metric["account__currency__prefix"],
|
||||||
|
"suffix": metric["account__currency__suffix"],
|
||||||
|
},
|
||||||
|
"final_total": final_total,
|
||||||
|
"income_total": metric["income_total"],
|
||||||
|
"expense_total": metric["expense_total"],
|
||||||
|
}
|
||||||
|
|
||||||
|
# Handle currency conversion if exchange currency is set
|
||||||
|
if metric["account__currency__exchange_currency"]:
|
||||||
|
from_currency = Currency.objects.get(id=currency_id)
|
||||||
|
exchange_currency = Currency.objects.get(
|
||||||
|
id=metric["account__currency__exchange_currency"]
|
||||||
|
)
|
||||||
|
|
||||||
|
converted_amount, prefix, suffix, decimal_places = convert(
|
||||||
|
amount=final_total,
|
||||||
|
from_currency=from_currency,
|
||||||
|
to_currency=exchange_currency,
|
||||||
|
)
|
||||||
|
|
||||||
|
if converted_amount is not None:
|
||||||
|
currency_data["exchanged"] = {
|
||||||
|
"final_total": converted_amount,
|
||||||
|
"currency": {
|
||||||
|
"prefix": prefix,
|
||||||
|
"suffix": suffix,
|
||||||
|
"decimal_places": decimal_places,
|
||||||
|
"code": exchange_currency.code,
|
||||||
|
"name": exchange_currency.name,
|
||||||
|
},
|
||||||
|
}
|
||||||
|
|
||||||
|
result["items"][item_key]["year_totals"][year]["currencies"][currency_id] = (
|
||||||
|
currency_data
|
||||||
|
)
|
||||||
|
|
||||||
|
# Accumulate item total (across all years for this item)
|
||||||
|
if currency_id not in result["items"][item_key]["total"]["currencies"]:
|
||||||
|
result["items"][item_key]["total"]["currencies"][currency_id] = {
|
||||||
|
"currency": currency_data["currency"].copy(),
|
||||||
|
"final_total": Decimal("0"),
|
||||||
|
}
|
||||||
|
result["items"][item_key]["total"]["currencies"][currency_id][
|
||||||
|
"final_total"
|
||||||
|
] += final_total
|
||||||
|
|
||||||
|
# Accumulate year total (across all items for this year)
|
||||||
|
if year not in result["year_totals"]:
|
||||||
|
result["year_totals"][year] = {"currencies": {}}
|
||||||
|
if currency_id not in result["year_totals"][year]["currencies"]:
|
||||||
|
result["year_totals"][year]["currencies"][currency_id] = {
|
||||||
|
"currency": currency_data["currency"].copy(),
|
||||||
|
"final_total": Decimal("0"),
|
||||||
|
}
|
||||||
|
result["year_totals"][year]["currencies"][currency_id]["final_total"] += (
|
||||||
|
final_total
|
||||||
|
)
|
||||||
|
|
||||||
|
# Accumulate grand total
|
||||||
|
if currency_id not in result["grand_total"]["currencies"]:
|
||||||
|
result["grand_total"]["currencies"][currency_id] = {
|
||||||
|
"currency": currency_data["currency"].copy(),
|
||||||
|
"final_total": Decimal("0"),
|
||||||
|
}
|
||||||
|
result["grand_total"]["currencies"][currency_id]["final_total"] += final_total
|
||||||
|
|
||||||
|
# Add currency conversion for item totals
|
||||||
|
for item_key, item_data in result["items"].items():
|
||||||
|
for currency_id, total_data in item_data["total"]["currencies"].items():
|
||||||
|
if currency_info[currency_id]["exchange_currency_id"]:
|
||||||
|
from_currency = Currency.objects.get(id=currency_id)
|
||||||
|
exchange_currency = Currency.objects.get(
|
||||||
|
id=currency_info[currency_id]["exchange_currency_id"]
|
||||||
|
)
|
||||||
|
converted_amount, prefix, suffix, decimal_places = convert(
|
||||||
|
amount=total_data["final_total"],
|
||||||
|
from_currency=from_currency,
|
||||||
|
to_currency=exchange_currency,
|
||||||
|
)
|
||||||
|
if converted_amount is not None:
|
||||||
|
total_data["exchanged"] = {
|
||||||
|
"final_total": converted_amount,
|
||||||
|
"currency": {
|
||||||
|
"prefix": prefix,
|
||||||
|
"suffix": suffix,
|
||||||
|
"decimal_places": decimal_places,
|
||||||
|
"code": exchange_currency.code,
|
||||||
|
"name": exchange_currency.name,
|
||||||
|
},
|
||||||
|
}
|
||||||
|
|
||||||
|
# Add currency conversion for year totals
|
||||||
|
for year, year_data in result["year_totals"].items():
|
||||||
|
for currency_id, total_data in year_data["currencies"].items():
|
||||||
|
if currency_info[currency_id]["exchange_currency_id"]:
|
||||||
|
from_currency = Currency.objects.get(id=currency_id)
|
||||||
|
exchange_currency = Currency.objects.get(
|
||||||
|
id=currency_info[currency_id]["exchange_currency_id"]
|
||||||
|
)
|
||||||
|
converted_amount, prefix, suffix, decimal_places = convert(
|
||||||
|
amount=total_data["final_total"],
|
||||||
|
from_currency=from_currency,
|
||||||
|
to_currency=exchange_currency,
|
||||||
|
)
|
||||||
|
if converted_amount is not None:
|
||||||
|
total_data["exchanged"] = {
|
||||||
|
"final_total": converted_amount,
|
||||||
|
"currency": {
|
||||||
|
"prefix": prefix,
|
||||||
|
"suffix": suffix,
|
||||||
|
"decimal_places": decimal_places,
|
||||||
|
"code": exchange_currency.code,
|
||||||
|
"name": exchange_currency.name,
|
||||||
|
},
|
||||||
|
}
|
||||||
|
|
||||||
|
# Add currency conversion for grand total
|
||||||
|
for currency_id, total_data in result["grand_total"]["currencies"].items():
|
||||||
|
if currency_info[currency_id]["exchange_currency_id"]:
|
||||||
|
from_currency = Currency.objects.get(id=currency_id)
|
||||||
|
exchange_currency = Currency.objects.get(
|
||||||
|
id=currency_info[currency_id]["exchange_currency_id"]
|
||||||
|
)
|
||||||
|
converted_amount, prefix, suffix, decimal_places = convert(
|
||||||
|
amount=total_data["final_total"],
|
||||||
|
from_currency=from_currency,
|
||||||
|
to_currency=exchange_currency,
|
||||||
|
)
|
||||||
|
if converted_amount is not None:
|
||||||
|
total_data["exchanged"] = {
|
||||||
|
"final_total": converted_amount,
|
||||||
|
"currency": {
|
||||||
|
"prefix": prefix,
|
||||||
|
"suffix": suffix,
|
||||||
|
"decimal_places": decimal_places,
|
||||||
|
"code": exchange_currency.code,
|
||||||
|
"name": exchange_currency.name,
|
||||||
|
},
|
||||||
|
}
|
||||||
|
|
||||||
|
return result
|
||||||
@@ -26,6 +26,8 @@ from apps.insights.utils.sankey import (
|
|||||||
generate_sankey_data_by_currency,
|
generate_sankey_data_by_currency,
|
||||||
)
|
)
|
||||||
from apps.insights.utils.transactions import get_transactions
|
from apps.insights.utils.transactions import get_transactions
|
||||||
|
from apps.insights.utils.year_by_year import get_year_by_year_data
|
||||||
|
from apps.insights.utils.month_by_month import get_month_by_month_data
|
||||||
from apps.transactions.models import TransactionCategory, Transaction
|
from apps.transactions.models import TransactionCategory, Transaction
|
||||||
from apps.transactions.utils.calculations import calculate_currency_totals
|
from apps.transactions.utils.calculations import calculate_currency_totals
|
||||||
|
|
||||||
@@ -306,3 +308,71 @@ def emergency_fund(request):
|
|||||||
"insights/fragments/emergency_fund.html",
|
"insights/fragments/emergency_fund.html",
|
||||||
{"data": currency_net_worth},
|
{"data": currency_net_worth},
|
||||||
)
|
)
|
||||||
|
|
||||||
|
|
||||||
|
@only_htmx
|
||||||
|
@login_required
|
||||||
|
@require_http_methods(["GET"])
|
||||||
|
def year_by_year(request):
|
||||||
|
if "group_by" in request.GET:
|
||||||
|
group_by = request.GET["group_by"]
|
||||||
|
request.session["insights_year_by_year_group_by"] = group_by
|
||||||
|
else:
|
||||||
|
group_by = request.session.get("insights_year_by_year_group_by", "categories")
|
||||||
|
|
||||||
|
# Validate group_by value
|
||||||
|
if group_by not in ("categories", "tags", "entities"):
|
||||||
|
group_by = "categories"
|
||||||
|
|
||||||
|
data = get_year_by_year_data(group_by=group_by)
|
||||||
|
|
||||||
|
return render(
|
||||||
|
request,
|
||||||
|
"insights/fragments/year_by_year.html",
|
||||||
|
{
|
||||||
|
"data": data,
|
||||||
|
"group_by": group_by,
|
||||||
|
},
|
||||||
|
)
|
||||||
|
|
||||||
|
|
||||||
|
@only_htmx
|
||||||
|
@login_required
|
||||||
|
@require_http_methods(["GET"])
|
||||||
|
def month_by_month(request):
|
||||||
|
# Handle year selection
|
||||||
|
if "year" in request.GET:
|
||||||
|
try:
|
||||||
|
year = int(request.GET["year"])
|
||||||
|
request.session["insights_month_by_month_year"] = year
|
||||||
|
except (ValueError, TypeError):
|
||||||
|
year = request.session.get(
|
||||||
|
"insights_month_by_month_year", timezone.localdate(timezone.now()).year
|
||||||
|
)
|
||||||
|
else:
|
||||||
|
year = request.session.get(
|
||||||
|
"insights_month_by_month_year", timezone.localdate(timezone.now()).year
|
||||||
|
)
|
||||||
|
|
||||||
|
# Handle group_by selection
|
||||||
|
if "group_by" in request.GET:
|
||||||
|
group_by = request.GET["group_by"]
|
||||||
|
request.session["insights_month_by_month_group_by"] = group_by
|
||||||
|
else:
|
||||||
|
group_by = request.session.get("insights_month_by_month_group_by", "categories")
|
||||||
|
|
||||||
|
# Validate group_by value
|
||||||
|
if group_by not in ("categories", "tags", "entities"):
|
||||||
|
group_by = "categories"
|
||||||
|
|
||||||
|
data = get_month_by_month_data(year=year, group_by=group_by)
|
||||||
|
|
||||||
|
return render(
|
||||||
|
request,
|
||||||
|
"insights/fragments/month_by_month.html",
|
||||||
|
{
|
||||||
|
"data": data,
|
||||||
|
"group_by": group_by,
|
||||||
|
"selected_year": year,
|
||||||
|
},
|
||||||
|
)
|
||||||
|
|||||||
@@ -0,0 +1,369 @@
|
|||||||
|
from datetime import date
|
||||||
|
from decimal import Decimal
|
||||||
|
|
||||||
|
from django.contrib.auth import get_user_model
|
||||||
|
from django.test import TestCase, override_settings
|
||||||
|
|
||||||
|
from apps.accounts.models import Account, AccountGroup
|
||||||
|
from apps.currencies.models import Currency
|
||||||
|
from apps.transactions.models import (
|
||||||
|
Transaction,
|
||||||
|
TransactionCategory,
|
||||||
|
TransactionTag,
|
||||||
|
)
|
||||||
|
|
||||||
|
|
||||||
|
@override_settings(
|
||||||
|
STORAGES={
|
||||||
|
"default": {"BACKEND": "django.core.files.storage.FileSystemStorage"},
|
||||||
|
"staticfiles": {
|
||||||
|
"BACKEND": "django.contrib.staticfiles.storage.StaticFilesStorage"
|
||||||
|
},
|
||||||
|
},
|
||||||
|
WHITENOISE_AUTOREFRESH=True,
|
||||||
|
)
|
||||||
|
class MonthlySummaryFilterBehaviorTests(TestCase):
|
||||||
|
"""Tests for monthly summary views filter behavior.
|
||||||
|
|
||||||
|
These tests verify that:
|
||||||
|
1. Views work correctly without any filters
|
||||||
|
2. Views work correctly with filters applied
|
||||||
|
3. The filter detection logic properly uses different querysets
|
||||||
|
4. Calculated values reflect the applied filters
|
||||||
|
"""
|
||||||
|
|
||||||
|
def setUp(self):
|
||||||
|
"""Set up test data"""
|
||||||
|
User = get_user_model()
|
||||||
|
self.user = User.objects.create_user(
|
||||||
|
email="testuser@test.com", password="testpass123"
|
||||||
|
)
|
||||||
|
self.client.login(username="testuser@test.com", password="testpass123")
|
||||||
|
|
||||||
|
self.currency = Currency.objects.create(
|
||||||
|
code="USD", name="US Dollar", decimal_places=2, prefix="$ "
|
||||||
|
)
|
||||||
|
self.account_group = AccountGroup.objects.create(name="Test Group")
|
||||||
|
self.account = Account.objects.create(
|
||||||
|
name="Test Account",
|
||||||
|
group=self.account_group,
|
||||||
|
currency=self.currency,
|
||||||
|
is_asset=False,
|
||||||
|
)
|
||||||
|
self.category = TransactionCategory.objects.create(
|
||||||
|
name="Test Category", owner=self.user
|
||||||
|
)
|
||||||
|
self.tag = TransactionTag.objects.create(name="TestTag", owner=self.user)
|
||||||
|
|
||||||
|
# Create test transactions for December 2025
|
||||||
|
# Income: 1000 (paid)
|
||||||
|
self.income_transaction = Transaction.objects.create(
|
||||||
|
account=self.account,
|
||||||
|
type=Transaction.Type.INCOME,
|
||||||
|
is_paid=True,
|
||||||
|
date=date(2025, 12, 10),
|
||||||
|
reference_date=date(2025, 12, 1),
|
||||||
|
amount=Decimal("1000.00"),
|
||||||
|
description="December Income",
|
||||||
|
owner=self.user,
|
||||||
|
)
|
||||||
|
|
||||||
|
# Expense: 200 (paid)
|
||||||
|
self.expense_transaction = Transaction.objects.create(
|
||||||
|
account=self.account,
|
||||||
|
type=Transaction.Type.EXPENSE,
|
||||||
|
is_paid=True,
|
||||||
|
date=date(2025, 12, 15),
|
||||||
|
reference_date=date(2025, 12, 1),
|
||||||
|
amount=Decimal("200.00"),
|
||||||
|
description="December Expense",
|
||||||
|
category=self.category,
|
||||||
|
owner=self.user,
|
||||||
|
)
|
||||||
|
self.expense_transaction.tags.add(self.tag)
|
||||||
|
|
||||||
|
# Expense: 150 (projected/unpaid)
|
||||||
|
self.projected_expense = Transaction.objects.create(
|
||||||
|
account=self.account,
|
||||||
|
type=Transaction.Type.EXPENSE,
|
||||||
|
is_paid=False,
|
||||||
|
date=date(2025, 12, 20),
|
||||||
|
reference_date=date(2025, 12, 1),
|
||||||
|
amount=Decimal("150.00"),
|
||||||
|
description="Projected Expense",
|
||||||
|
owner=self.user,
|
||||||
|
)
|
||||||
|
|
||||||
|
def _get_currency_data(self, context_dict):
|
||||||
|
"""Helper to extract data for our test currency from context dict.
|
||||||
|
|
||||||
|
The context dict is keyed by currency ID, so we need to find
|
||||||
|
the entry for our currency.
|
||||||
|
"""
|
||||||
|
if not context_dict:
|
||||||
|
return None
|
||||||
|
for currency_id, data in context_dict.items():
|
||||||
|
if data.get("currency", {}).get("code") == "USD":
|
||||||
|
return data
|
||||||
|
return None
|
||||||
|
|
||||||
|
def _create_asset_income(self):
|
||||||
|
asset_account = Account.objects.create(
|
||||||
|
name="Asset Account",
|
||||||
|
group=self.account_group,
|
||||||
|
currency=self.currency,
|
||||||
|
is_asset=True,
|
||||||
|
)
|
||||||
|
Transaction.objects.create(
|
||||||
|
account=asset_account,
|
||||||
|
type=Transaction.Type.INCOME,
|
||||||
|
is_paid=True,
|
||||||
|
date=date(2025, 12, 25),
|
||||||
|
reference_date=date(2025, 12, 1),
|
||||||
|
amount=Decimal("50.00"),
|
||||||
|
description="Asset Income",
|
||||||
|
owner=self.user,
|
||||||
|
)
|
||||||
|
return asset_account
|
||||||
|
|
||||||
|
# --- monthly_summary view tests ---
|
||||||
|
|
||||||
|
def test_monthly_summary_no_filter_returns_200(self):
|
||||||
|
"""Test that monthly_summary returns 200 without filters"""
|
||||||
|
response = self.client.get(
|
||||||
|
"/monthly/12/2025/summary/",
|
||||||
|
HTTP_HX_REQUEST="true",
|
||||||
|
)
|
||||||
|
self.assertEqual(response.status_code, 200)
|
||||||
|
|
||||||
|
def test_monthly_summary_no_filter_includes_all_transactions(self):
|
||||||
|
"""Without filters, summary should include all transactions"""
|
||||||
|
response = self.client.get(
|
||||||
|
"/monthly/12/2025/summary/",
|
||||||
|
HTTP_HX_REQUEST="true",
|
||||||
|
)
|
||||||
|
context = response.context
|
||||||
|
|
||||||
|
# income_current should have the income: 1000
|
||||||
|
income_current = context.get("income_current", {})
|
||||||
|
usd_data = self._get_currency_data(income_current)
|
||||||
|
self.assertIsNotNone(usd_data)
|
||||||
|
self.assertEqual(usd_data["income_current"], Decimal("1000.00"))
|
||||||
|
|
||||||
|
# expense_current should have paid expense: 200
|
||||||
|
expense_current = context.get("expense_current", {})
|
||||||
|
usd_data = self._get_currency_data(expense_current)
|
||||||
|
self.assertIsNotNone(usd_data)
|
||||||
|
self.assertEqual(usd_data["expense_current"], Decimal("200.00"))
|
||||||
|
|
||||||
|
# expense_projected should have unpaid expense: 150
|
||||||
|
expense_projected = context.get("expense_projected", {})
|
||||||
|
usd_data = self._get_currency_data(expense_projected)
|
||||||
|
self.assertIsNotNone(usd_data)
|
||||||
|
self.assertEqual(usd_data["expense_projected"], Decimal("150.00"))
|
||||||
|
|
||||||
|
def test_monthly_summary_type_filter_only_income(self):
|
||||||
|
"""With type=IN filter, summary should only include income"""
|
||||||
|
response = self.client.get(
|
||||||
|
"/monthly/12/2025/summary/?type=IN",
|
||||||
|
HTTP_HX_REQUEST="true",
|
||||||
|
)
|
||||||
|
context = response.context
|
||||||
|
|
||||||
|
# income_current should still have 1000
|
||||||
|
income_current = context.get("income_current", {})
|
||||||
|
usd_data = self._get_currency_data(income_current)
|
||||||
|
self.assertIsNotNone(usd_data)
|
||||||
|
self.assertEqual(usd_data["income_current"], Decimal("1000.00"))
|
||||||
|
|
||||||
|
# expense_current should be empty/zero (filtered out)
|
||||||
|
expense_current = context.get("expense_current", {})
|
||||||
|
usd_data = self._get_currency_data(expense_current)
|
||||||
|
if usd_data:
|
||||||
|
self.assertEqual(usd_data.get("expense_current", 0), Decimal("0"))
|
||||||
|
|
||||||
|
# expense_projected should be empty/zero (filtered out)
|
||||||
|
expense_projected = context.get("expense_projected", {})
|
||||||
|
usd_data = self._get_currency_data(expense_projected)
|
||||||
|
if usd_data:
|
||||||
|
self.assertEqual(usd_data.get("expense_projected", 0), Decimal("0"))
|
||||||
|
|
||||||
|
def test_monthly_summary_type_filter_only_expenses(self):
|
||||||
|
"""With type=EX filter, summary should only include expenses"""
|
||||||
|
response = self.client.get(
|
||||||
|
"/monthly/12/2025/summary/?type=EX",
|
||||||
|
HTTP_HX_REQUEST="true",
|
||||||
|
)
|
||||||
|
context = response.context
|
||||||
|
|
||||||
|
# income_current should be empty/zero (filtered out)
|
||||||
|
income_current = context.get("income_current", {})
|
||||||
|
usd_data = self._get_currency_data(income_current)
|
||||||
|
if usd_data:
|
||||||
|
self.assertEqual(usd_data.get("income_current", 0), Decimal("0"))
|
||||||
|
|
||||||
|
# expense_current should have 200
|
||||||
|
expense_current = context.get("expense_current", {})
|
||||||
|
usd_data = self._get_currency_data(expense_current)
|
||||||
|
self.assertIsNotNone(usd_data)
|
||||||
|
self.assertEqual(usd_data["expense_current"], Decimal("200.00"))
|
||||||
|
|
||||||
|
# expense_projected should have 150
|
||||||
|
expense_projected = context.get("expense_projected", {})
|
||||||
|
usd_data = self._get_currency_data(expense_projected)
|
||||||
|
self.assertIsNotNone(usd_data)
|
||||||
|
self.assertEqual(usd_data["expense_projected"], Decimal("150.00"))
|
||||||
|
|
||||||
|
def test_monthly_summary_is_paid_filter_only_paid(self):
|
||||||
|
"""With is_paid=1 filter, summary should only include paid transactions"""
|
||||||
|
response = self.client.get(
|
||||||
|
"/monthly/12/2025/summary/?is_paid=1",
|
||||||
|
HTTP_HX_REQUEST="true",
|
||||||
|
)
|
||||||
|
context = response.context
|
||||||
|
|
||||||
|
# income_current should have 1000 (paid)
|
||||||
|
income_current = context.get("income_current", {})
|
||||||
|
usd_data = self._get_currency_data(income_current)
|
||||||
|
self.assertIsNotNone(usd_data)
|
||||||
|
self.assertEqual(usd_data["income_current"], Decimal("1000.00"))
|
||||||
|
|
||||||
|
# expense_current should have 200 (paid)
|
||||||
|
expense_current = context.get("expense_current", {})
|
||||||
|
usd_data = self._get_currency_data(expense_current)
|
||||||
|
self.assertIsNotNone(usd_data)
|
||||||
|
self.assertEqual(usd_data["expense_current"], Decimal("200.00"))
|
||||||
|
|
||||||
|
# expense_projected should be empty/zero (filtered out - unpaid)
|
||||||
|
expense_projected = context.get("expense_projected", {})
|
||||||
|
usd_data = self._get_currency_data(expense_projected)
|
||||||
|
if usd_data:
|
||||||
|
self.assertEqual(usd_data.get("expense_projected", 0), Decimal("0"))
|
||||||
|
|
||||||
|
def test_monthly_summary_is_paid_filter_only_unpaid(self):
|
||||||
|
"""With is_paid=0 filter, summary should only include unpaid transactions"""
|
||||||
|
response = self.client.get(
|
||||||
|
"/monthly/12/2025/summary/?is_paid=0",
|
||||||
|
HTTP_HX_REQUEST="true",
|
||||||
|
)
|
||||||
|
context = response.context
|
||||||
|
|
||||||
|
# income_current should be empty/zero (filtered out - paid)
|
||||||
|
income_current = context.get("income_current", {})
|
||||||
|
usd_data = self._get_currency_data(income_current)
|
||||||
|
if usd_data:
|
||||||
|
self.assertEqual(usd_data.get("income_current", 0), Decimal("0"))
|
||||||
|
|
||||||
|
# expense_current should be empty/zero (filtered out - paid)
|
||||||
|
expense_current = context.get("expense_current", {})
|
||||||
|
usd_data = self._get_currency_data(expense_current)
|
||||||
|
if usd_data:
|
||||||
|
self.assertEqual(usd_data.get("expense_current", 0), Decimal("0"))
|
||||||
|
|
||||||
|
# expense_projected should have 150 (unpaid)
|
||||||
|
expense_projected = context.get("expense_projected", {})
|
||||||
|
usd_data = self._get_currency_data(expense_projected)
|
||||||
|
self.assertIsNotNone(usd_data)
|
||||||
|
self.assertEqual(usd_data["expense_projected"], Decimal("150.00"))
|
||||||
|
|
||||||
|
def test_monthly_summary_description_filter(self):
|
||||||
|
"""With description filter, summary should only include matching transactions"""
|
||||||
|
response = self.client.get(
|
||||||
|
"/monthly/12/2025/summary/?description=Income",
|
||||||
|
HTTP_HX_REQUEST="true",
|
||||||
|
)
|
||||||
|
context = response.context
|
||||||
|
|
||||||
|
# Only income matches "Income" description
|
||||||
|
income_current = context.get("income_current", {})
|
||||||
|
usd_data = self._get_currency_data(income_current)
|
||||||
|
self.assertIsNotNone(usd_data)
|
||||||
|
self.assertEqual(usd_data["income_current"], Decimal("1000.00"))
|
||||||
|
|
||||||
|
# Expenses should be filtered out
|
||||||
|
expense_current = context.get("expense_current", {})
|
||||||
|
usd_data = self._get_currency_data(expense_current)
|
||||||
|
if usd_data:
|
||||||
|
self.assertEqual(usd_data.get("expense_current", 0), Decimal("0"))
|
||||||
|
|
||||||
|
def test_monthly_summary_amount_filter(self):
|
||||||
|
"""With amount filter, summary should only include transactions in range"""
|
||||||
|
# Filter to only get transactions between 100 and 250 (should get 200 and 150)
|
||||||
|
response = self.client.get(
|
||||||
|
"/monthly/12/2025/summary/?from_amount=100&to_amount=250",
|
||||||
|
HTTP_HX_REQUEST="true",
|
||||||
|
)
|
||||||
|
context = response.context
|
||||||
|
|
||||||
|
# Income (1000) should be filtered out
|
||||||
|
income_current = context.get("income_current", {})
|
||||||
|
usd_data = self._get_currency_data(income_current)
|
||||||
|
if usd_data:
|
||||||
|
self.assertEqual(usd_data.get("income_current", 0), Decimal("0"))
|
||||||
|
|
||||||
|
# expense_current should have 200
|
||||||
|
expense_current = context.get("expense_current", {})
|
||||||
|
usd_data = self._get_currency_data(expense_current)
|
||||||
|
self.assertIsNotNone(usd_data)
|
||||||
|
self.assertEqual(usd_data["expense_current"], Decimal("200.00"))
|
||||||
|
|
||||||
|
# expense_projected should have 150
|
||||||
|
expense_projected = context.get("expense_projected", {})
|
||||||
|
usd_data = self._get_currency_data(expense_projected)
|
||||||
|
self.assertIsNotNone(usd_data)
|
||||||
|
self.assertEqual(usd_data["expense_projected"], Decimal("150.00"))
|
||||||
|
|
||||||
|
# --- monthly_account_summary view tests ---
|
||||||
|
|
||||||
|
def test_monthly_account_summary_no_filter_returns_200(self):
|
||||||
|
"""Test that monthly_account_summary returns 200 without filters"""
|
||||||
|
response = self.client.get(
|
||||||
|
"/monthly/12/2025/summary/accounts/",
|
||||||
|
HTTP_HX_REQUEST="true",
|
||||||
|
)
|
||||||
|
self.assertEqual(response.status_code, 200)
|
||||||
|
|
||||||
|
def test_monthly_account_summary_includes_asset_accounts(self):
|
||||||
|
asset_account = self._create_asset_income()
|
||||||
|
response = self.client.get(
|
||||||
|
"/monthly/12/2025/summary/accounts/",
|
||||||
|
HTTP_HX_REQUEST="true",
|
||||||
|
)
|
||||||
|
|
||||||
|
self.assertIn(asset_account.id, response.context["account_data"])
|
||||||
|
|
||||||
|
def test_monthly_account_summary_with_filter_returns_200(self):
|
||||||
|
"""Test that monthly_account_summary returns 200 with filter"""
|
||||||
|
response = self.client.get(
|
||||||
|
"/monthly/12/2025/summary/accounts/?type=IN",
|
||||||
|
HTTP_HX_REQUEST="true",
|
||||||
|
)
|
||||||
|
self.assertEqual(response.status_code, 200)
|
||||||
|
|
||||||
|
# --- monthly_currency_summary view tests ---
|
||||||
|
|
||||||
|
def test_monthly_currency_summary_no_filter_returns_200(self):
|
||||||
|
"""Test that monthly_currency_summary returns 200 without filters"""
|
||||||
|
response = self.client.get(
|
||||||
|
"/monthly/12/2025/summary/currencies/",
|
||||||
|
HTTP_HX_REQUEST="true",
|
||||||
|
)
|
||||||
|
self.assertEqual(response.status_code, 200)
|
||||||
|
|
||||||
|
def test_monthly_currency_summary_includes_asset_account_transactions(self):
|
||||||
|
self._create_asset_income()
|
||||||
|
response = self.client.get(
|
||||||
|
"/monthly/12/2025/summary/currencies/",
|
||||||
|
HTTP_HX_REQUEST="true",
|
||||||
|
)
|
||||||
|
|
||||||
|
usd_data = self._get_currency_data(response.context["currency_data"])
|
||||||
|
self.assertEqual(usd_data["income_current"], Decimal("1050.00"))
|
||||||
|
|
||||||
|
def test_monthly_currency_summary_with_filter_returns_200(self):
|
||||||
|
"""Test that monthly_currency_summary returns 200 with filter"""
|
||||||
|
response = self.client.get(
|
||||||
|
"/monthly/12/2025/summary/currencies/?type=EX",
|
||||||
|
HTTP_HX_REQUEST="true",
|
||||||
|
)
|
||||||
|
self.assertEqual(response.status_code, 200)
|
||||||
@@ -2,7 +2,8 @@ from django.contrib.auth.decorators import login_required
|
|||||||
from django.db.models import (
|
from django.db.models import (
|
||||||
Q,
|
Q,
|
||||||
)
|
)
|
||||||
from django.http import HttpResponse
|
from django.http import HttpResponse, Http404
|
||||||
|
|
||||||
from django.shortcuts import render, redirect
|
from django.shortcuts import render, redirect
|
||||||
from django.utils import timezone
|
from django.utils import timezone
|
||||||
from django.views.decorators.http import require_http_methods
|
from django.views.decorators.http import require_http_methods
|
||||||
@@ -13,7 +14,7 @@ from apps.monthly_overview.utils.daily_spending_allowance import (
|
|||||||
calculate_daily_allowance_currency,
|
calculate_daily_allowance_currency,
|
||||||
)
|
)
|
||||||
from apps.transactions.filters import TransactionsFilter
|
from apps.transactions.filters import TransactionsFilter
|
||||||
from apps.transactions.models import Transaction
|
from apps.transactions.models import FilterPreset, Transaction
|
||||||
from apps.transactions.utils.calculations import (
|
from apps.transactions.utils.calculations import (
|
||||||
calculate_currency_totals,
|
calculate_currency_totals,
|
||||||
calculate_percentage_distribution,
|
calculate_percentage_distribution,
|
||||||
@@ -36,8 +37,6 @@ def monthly_overview(request, month: int, year: int):
|
|||||||
summary_tab = request.session.get("monthly_summary_tab", "summary")
|
summary_tab = request.session.get("monthly_summary_tab", "summary")
|
||||||
|
|
||||||
if month < 1 or month > 12:
|
if month < 1 or month > 12:
|
||||||
from django.http import Http404
|
|
||||||
|
|
||||||
raise Http404("Month is out of range")
|
raise Http404("Month is out of range")
|
||||||
|
|
||||||
next_month = 1 if month == 12 else month + 1
|
next_month = 1 if month == 12 else month + 1
|
||||||
@@ -59,6 +58,8 @@ def monthly_overview(request, month: int, year: int):
|
|||||||
"previous_month": previous_month,
|
"previous_month": previous_month,
|
||||||
"previous_year": previous_year,
|
"previous_year": previous_year,
|
||||||
"filter": f,
|
"filter": f,
|
||||||
|
"filter_is_active": f.has_active_filters,
|
||||||
|
"filter_presets": FilterPreset.objects.filter(owner=request.user),
|
||||||
"order": order,
|
"order": order,
|
||||||
"summary_tab": summary_tab,
|
"summary_tab": summary_tab,
|
||||||
},
|
},
|
||||||
@@ -76,6 +77,8 @@ def transactions_list(request, month: int, year: int):
|
|||||||
if order != request.session.get("monthly_transactions_order", "default"):
|
if order != request.session.get("monthly_transactions_order", "default"):
|
||||||
request.session["monthly_transactions_order"] = order
|
request.session["monthly_transactions_order"] = order
|
||||||
|
|
||||||
|
today = timezone.localdate(timezone.now())
|
||||||
|
|
||||||
f = TransactionsFilter(request.GET)
|
f = TransactionsFilter(request.GET)
|
||||||
transactions_filtered = f.qs.filter(
|
transactions_filtered = f.qs.filter(
|
||||||
reference_date__year=year,
|
reference_date__year=year,
|
||||||
@@ -93,12 +96,28 @@ def transactions_list(request, month: int, year: int):
|
|||||||
"dca_income_entries",
|
"dca_income_entries",
|
||||||
)
|
)
|
||||||
|
|
||||||
|
# Late transactions: date < today and is_paid = False (only shown for default ordering)
|
||||||
|
late_transactions = None
|
||||||
|
if order == "default":
|
||||||
|
late_transactions = transactions_filtered.filter(
|
||||||
|
date__lt=today,
|
||||||
|
is_paid=False,
|
||||||
|
).order_by("date", "id")
|
||||||
|
# Exclude late transactions from the main list
|
||||||
|
transactions_filtered = transactions_filtered.exclude(
|
||||||
|
date__lt=today,
|
||||||
|
is_paid=False,
|
||||||
|
)
|
||||||
|
|
||||||
transactions_filtered = default_order(transactions_filtered, order=order)
|
transactions_filtered = default_order(transactions_filtered, order=order)
|
||||||
|
|
||||||
return render(
|
return render(
|
||||||
request,
|
request,
|
||||||
"monthly_overview/fragments/list.html",
|
"monthly_overview/fragments/list.html",
|
||||||
context={"transactions": transactions_filtered},
|
context={
|
||||||
|
"transactions": transactions_filtered,
|
||||||
|
"late_transactions": late_transactions,
|
||||||
|
},
|
||||||
)
|
)
|
||||||
|
|
||||||
|
|
||||||
@@ -107,17 +126,48 @@ def transactions_list(request, month: int, year: int):
|
|||||||
@require_http_methods(["GET"])
|
@require_http_methods(["GET"])
|
||||||
def monthly_summary(request, month: int, year: int):
|
def monthly_summary(request, month: int, year: int):
|
||||||
# Base queryset with all required filters
|
# Base queryset with all required filters
|
||||||
base_queryset = (
|
base_queryset = Transaction.objects.filter(
|
||||||
Transaction.objects.filter(
|
reference_date__year=year,
|
||||||
reference_date__year=year,
|
reference_date__month=month,
|
||||||
reference_date__month=month,
|
|
||||||
account__is_asset=False,
|
|
||||||
)
|
|
||||||
.exclude(Q(Q(category__mute=True) & ~Q(category=None)) | Q(mute=True))
|
|
||||||
.exclude(account__in=request.user.untracked_accounts.all())
|
|
||||||
)
|
)
|
||||||
|
|
||||||
data = calculate_currency_totals(base_queryset, ignore_empty=True)
|
# Apply filters and check if any are active
|
||||||
|
f = TransactionsFilter(request.GET, queryset=base_queryset)
|
||||||
|
|
||||||
|
# Check if any filter has a non-default value
|
||||||
|
# Default values are: type=['IN', 'EX'], is_paid=['1', '0'], everything else empty
|
||||||
|
has_active_filter = False
|
||||||
|
if f.form.is_valid():
|
||||||
|
for name, value in f.form.cleaned_data.items():
|
||||||
|
# Skip fields with default/empty values
|
||||||
|
if not value:
|
||||||
|
continue
|
||||||
|
# Skip type if it has both default values
|
||||||
|
if name == "type" and set(value) == {"IN", "EX"}:
|
||||||
|
continue
|
||||||
|
# Skip is_paid if it has both default values (values are strings)
|
||||||
|
if name == "is_paid" and set(value) == {"1", "0"}:
|
||||||
|
continue
|
||||||
|
# Skip mute_status if it has both default values
|
||||||
|
if name == "mute_status" and set(value) == {"active", "muted"}:
|
||||||
|
continue
|
||||||
|
# If we get here, there's an active filter
|
||||||
|
has_active_filter = True
|
||||||
|
break
|
||||||
|
|
||||||
|
if has_active_filter:
|
||||||
|
queryset = f.qs
|
||||||
|
else:
|
||||||
|
queryset = (
|
||||||
|
base_queryset.exclude(
|
||||||
|
Q(Q(category__mute=True) & ~Q(category=None)) | Q(mute=True)
|
||||||
|
)
|
||||||
|
.exclude(account__in=request.user.untracked_accounts.all())
|
||||||
|
.exclude(account__is_asset=True)
|
||||||
|
)
|
||||||
|
|
||||||
|
data = calculate_currency_totals(queryset, ignore_empty=True)
|
||||||
|
|
||||||
percentages = calculate_percentage_distribution(data)
|
percentages = calculate_percentage_distribution(data)
|
||||||
|
|
||||||
context = {
|
context = {
|
||||||
@@ -132,6 +182,7 @@ def monthly_summary(request, month: int, year: int):
|
|||||||
currency_totals=data, month=month, year=year
|
currency_totals=data, month=month, year=year
|
||||||
),
|
),
|
||||||
"percentages": percentages,
|
"percentages": percentages,
|
||||||
|
"has_active_filter": has_active_filter,
|
||||||
}
|
}
|
||||||
|
|
||||||
return render(
|
return render(
|
||||||
@@ -149,9 +200,37 @@ def monthly_account_summary(request, month: int, year: int):
|
|||||||
base_queryset = Transaction.objects.filter(
|
base_queryset = Transaction.objects.filter(
|
||||||
reference_date__year=year,
|
reference_date__year=year,
|
||||||
reference_date__month=month,
|
reference_date__month=month,
|
||||||
).exclude(Q(Q(category__mute=True) & ~Q(category=None)) | Q(mute=True))
|
)
|
||||||
|
|
||||||
account_data = calculate_account_totals(transactions_queryset=base_queryset.all())
|
# Apply filters and check if any are active
|
||||||
|
f = TransactionsFilter(request.GET, queryset=base_queryset)
|
||||||
|
|
||||||
|
# Check if any filter has a non-default value
|
||||||
|
has_active_filter = False
|
||||||
|
if f.form.is_valid():
|
||||||
|
for name, value in f.form.cleaned_data.items():
|
||||||
|
if not value:
|
||||||
|
continue
|
||||||
|
if name == "type" and set(value) == {"IN", "EX"}:
|
||||||
|
continue
|
||||||
|
if name == "is_paid" and set(value) == {"1", "0"}:
|
||||||
|
continue
|
||||||
|
if name == "mute_status" and set(value) == {"active", "muted"}:
|
||||||
|
continue
|
||||||
|
has_active_filter = True
|
||||||
|
break
|
||||||
|
|
||||||
|
if has_active_filter:
|
||||||
|
queryset = f.qs
|
||||||
|
else:
|
||||||
|
queryset = (
|
||||||
|
base_queryset.exclude(
|
||||||
|
Q(Q(category__mute=True) & ~Q(category=None)) | Q(mute=True)
|
||||||
|
)
|
||||||
|
.exclude(account__in=request.user.untracked_accounts.all())
|
||||||
|
)
|
||||||
|
|
||||||
|
account_data = calculate_account_totals(transactions_queryset=queryset.all())
|
||||||
account_percentages = calculate_percentage_distribution(account_data)
|
account_percentages = calculate_percentage_distribution(account_data)
|
||||||
|
|
||||||
context = {
|
context = {
|
||||||
@@ -171,16 +250,40 @@ def monthly_account_summary(request, month: int, year: int):
|
|||||||
@require_http_methods(["GET"])
|
@require_http_methods(["GET"])
|
||||||
def monthly_currency_summary(request, month: int, year: int):
|
def monthly_currency_summary(request, month: int, year: int):
|
||||||
# Base queryset with all required filters
|
# Base queryset with all required filters
|
||||||
base_queryset = (
|
base_queryset = Transaction.objects.filter(
|
||||||
Transaction.objects.filter(
|
reference_date__year=year,
|
||||||
reference_date__year=year,
|
reference_date__month=month,
|
||||||
reference_date__month=month,
|
|
||||||
)
|
|
||||||
.exclude(Q(Q(category__mute=True) & ~Q(category=None)) | Q(mute=True))
|
|
||||||
.exclude(account__in=request.user.untracked_accounts.all())
|
|
||||||
)
|
)
|
||||||
|
|
||||||
currency_data = calculate_currency_totals(base_queryset.all(), ignore_empty=True)
|
# Apply filters and check if any are active
|
||||||
|
f = TransactionsFilter(request.GET, queryset=base_queryset)
|
||||||
|
|
||||||
|
# Check if any filter has a non-default value
|
||||||
|
has_active_filter = False
|
||||||
|
if f.form.is_valid():
|
||||||
|
for name, value in f.form.cleaned_data.items():
|
||||||
|
if not value:
|
||||||
|
continue
|
||||||
|
if name == "type" and set(value) == {"IN", "EX"}:
|
||||||
|
continue
|
||||||
|
if name == "is_paid" and set(value) == {"1", "0"}:
|
||||||
|
continue
|
||||||
|
if name == "mute_status" and set(value) == {"active", "muted"}:
|
||||||
|
continue
|
||||||
|
has_active_filter = True
|
||||||
|
break
|
||||||
|
|
||||||
|
if has_active_filter:
|
||||||
|
queryset = f.qs
|
||||||
|
else:
|
||||||
|
queryset = (
|
||||||
|
base_queryset.exclude(
|
||||||
|
Q(Q(category__mute=True) & ~Q(category=None)) | Q(mute=True)
|
||||||
|
)
|
||||||
|
.exclude(account__in=request.user.untracked_accounts.all())
|
||||||
|
)
|
||||||
|
|
||||||
|
currency_data = calculate_currency_totals(queryset.all(), ignore_empty=True)
|
||||||
currency_percentages = calculate_percentage_distribution(currency_data)
|
currency_percentages = calculate_percentage_distribution(currency_data)
|
||||||
|
|
||||||
context = {
|
context = {
|
||||||
|
|||||||
@@ -1,3 +1,82 @@
|
|||||||
from django.test import TestCase
|
import json
|
||||||
|
import tempfile
|
||||||
|
from datetime import date
|
||||||
|
from decimal import Decimal
|
||||||
|
|
||||||
# Create your tests here.
|
from django.contrib.auth import get_user_model
|
||||||
|
from django.test import TestCase, override_settings
|
||||||
|
from django.urls import reverse
|
||||||
|
from django.utils import timezone
|
||||||
|
|
||||||
|
from apps.accounts.models import Account
|
||||||
|
from apps.currencies.models import Currency, ExchangeRate
|
||||||
|
from apps.transactions.models import Transaction
|
||||||
|
|
||||||
|
|
||||||
|
@override_settings(
|
||||||
|
STATIC_ROOT=tempfile.gettempdir(),
|
||||||
|
STORAGES={
|
||||||
|
"default": {"BACKEND": "django.core.files.storage.FileSystemStorage"},
|
||||||
|
"staticfiles": {
|
||||||
|
"BACKEND": "django.contrib.staticfiles.storage.StaticFilesStorage"
|
||||||
|
},
|
||||||
|
},
|
||||||
|
)
|
||||||
|
class NetWorthCurrencyChartTests(TestCase):
|
||||||
|
def test_consolidated_currency_is_a_selectable_dashed_matching_color_line(self):
|
||||||
|
user = get_user_model().objects.create_user(
|
||||||
|
email="chart@example.com", password="password"
|
||||||
|
)
|
||||||
|
usd = Currency.objects.create(code="USD", name="US Dollar", prefix="$ ")
|
||||||
|
eur = Currency.objects.create(
|
||||||
|
code="EUR", name="Euro", prefix="€ ", exchange_currency=usd
|
||||||
|
)
|
||||||
|
usd_account = Account.all_objects.create(
|
||||||
|
name="USD account", currency=usd, owner=user
|
||||||
|
)
|
||||||
|
eur_account = Account.all_objects.create(
|
||||||
|
name="EUR account", currency=eur, owner=user
|
||||||
|
)
|
||||||
|
ExchangeRate.objects.create(
|
||||||
|
from_currency=eur,
|
||||||
|
to_currency=usd,
|
||||||
|
rate=Decimal("1.234567"),
|
||||||
|
date=timezone.now(),
|
||||||
|
)
|
||||||
|
for account, amount in ((usd_account, "100"), (eur_account, "50")):
|
||||||
|
Transaction.userless_all_objects.create(
|
||||||
|
account=account,
|
||||||
|
owner=user,
|
||||||
|
type=Transaction.Type.INCOME,
|
||||||
|
amount=Decimal(amount),
|
||||||
|
date=date(2026, 1, 15),
|
||||||
|
reference_date=date(2026, 1, 1),
|
||||||
|
is_paid=True,
|
||||||
|
)
|
||||||
|
|
||||||
|
self.client.force_login(user)
|
||||||
|
response = self.client.get(reverse("net_worth"))
|
||||||
|
|
||||||
|
self.assertEqual(response.status_code, 200)
|
||||||
|
chart_data = json.loads(response.context["chart_data_currency_json"])
|
||||||
|
datasets = {dataset["label"]: dataset for dataset in chart_data["datasets"]}
|
||||||
|
self.assertIn("US Dollar Consolidated", datasets)
|
||||||
|
regular = datasets["US Dollar"]
|
||||||
|
consolidated = datasets["US Dollar Consolidated"]
|
||||||
|
self.assertEqual(consolidated["data"], [161.73])
|
||||||
|
self.assertNotIn("borderColor", regular)
|
||||||
|
self.assertNotIn("borderColor", consolidated)
|
||||||
|
self.assertEqual(consolidated["colorSource"], "US Dollar")
|
||||||
|
self.assertEqual(consolidated["borderDash"], [12, 6])
|
||||||
|
self.assertEqual(consolidated["pointRadius"], 0)
|
||||||
|
self.assertEqual(consolidated["pointHitRadius"], 8)
|
||||||
|
self.assertContains(
|
||||||
|
response,
|
||||||
|
"showOnlyCurrencyDataset('US Dollar Consolidated', 'US Dollar')",
|
||||||
|
html=False,
|
||||||
|
)
|
||||||
|
self.assertContains(
|
||||||
|
response,
|
||||||
|
'<span class="text-start shrink">Consolidated</span>',
|
||||||
|
html=False,
|
||||||
|
)
|
||||||
|
|||||||
@@ -3,8 +3,11 @@ import json
|
|||||||
from django.contrib.auth.decorators import login_required
|
from django.contrib.auth.decorators import login_required
|
||||||
from django.core.serializers.json import DjangoJSONEncoder
|
from django.core.serializers.json import DjangoJSONEncoder
|
||||||
from django.shortcuts import render, redirect
|
from django.shortcuts import render, redirect
|
||||||
|
from django.utils.translation import gettext
|
||||||
from django.views.decorators.http import require_http_methods
|
from django.views.decorators.http import require_http_methods
|
||||||
|
|
||||||
|
from apps.currencies.models import Currency
|
||||||
|
from apps.currencies.utils.convert import convert
|
||||||
from apps.net_worth.utils.calculate_net_worth import (
|
from apps.net_worth.utils.calculate_net_worth import (
|
||||||
calculate_historical_currency_net_worth,
|
calculate_historical_currency_net_worth,
|
||||||
calculate_historical_account_balance,
|
calculate_historical_account_balance,
|
||||||
@@ -78,6 +81,17 @@ def net_worth(request):
|
|||||||
)
|
)
|
||||||
|
|
||||||
datasets = []
|
datasets = []
|
||||||
|
currency_models = {
|
||||||
|
currency.name: currency
|
||||||
|
for currency in Currency.objects.filter(name__in=currencies).select_related(
|
||||||
|
"exchange_currency"
|
||||||
|
)
|
||||||
|
}
|
||||||
|
consolidated_currencies = {
|
||||||
|
data["currency"]["name"]
|
||||||
|
for data in currency_net_worth.values()
|
||||||
|
if data["consolidated"]["total_final"] != data["total_final"]
|
||||||
|
}
|
||||||
for i, currency in enumerate(currencies):
|
for i, currency in enumerate(currencies):
|
||||||
data = [
|
data = [
|
||||||
float(month_data[currency])
|
float(month_data[currency])
|
||||||
@@ -93,6 +107,45 @@ def net_worth(request):
|
|||||||
}
|
}
|
||||||
)
|
)
|
||||||
|
|
||||||
|
if currency in consolidated_currencies:
|
||||||
|
target = currency_models[currency]
|
||||||
|
sources = [
|
||||||
|
source
|
||||||
|
for source in currency_models.values()
|
||||||
|
if source.exchange_currency_id == target.id
|
||||||
|
]
|
||||||
|
rates = {}
|
||||||
|
for source in sources:
|
||||||
|
converted, _, _, _ = convert(1, source, target)
|
||||||
|
if converted is not None:
|
||||||
|
rates[source.name] = converted
|
||||||
|
|
||||||
|
consolidated_data = [
|
||||||
|
float(
|
||||||
|
round(
|
||||||
|
month_data[currency]
|
||||||
|
+ sum(
|
||||||
|
month_data[source] * rate for source, rate in rates.items()
|
||||||
|
),
|
||||||
|
target.decimal_places,
|
||||||
|
)
|
||||||
|
)
|
||||||
|
for month_data in historical_currency_net_worth.values()
|
||||||
|
]
|
||||||
|
datasets.append(
|
||||||
|
{
|
||||||
|
"label": f"{currency} {gettext('Consolidated')}",
|
||||||
|
"data": consolidated_data,
|
||||||
|
"yAxisID": f"y{i}",
|
||||||
|
"fill": False,
|
||||||
|
"tension": 0.1,
|
||||||
|
"colorSource": currency,
|
||||||
|
"borderDash": [12, 6],
|
||||||
|
"pointRadius": 0,
|
||||||
|
"pointHitRadius": 8,
|
||||||
|
}
|
||||||
|
)
|
||||||
|
|
||||||
chart_data_currency = {"labels": labels, "datasets": datasets}
|
chart_data_currency = {"labels": labels, "datasets": datasets}
|
||||||
|
|
||||||
chart_data_currency_json = json.dumps(chart_data_currency, cls=DjangoJSONEncoder)
|
chart_data_currency_json = json.dumps(chart_data_currency, cls=DjangoJSONEncoder)
|
||||||
|
|||||||
@@ -365,7 +365,9 @@ def check_for_transaction_rules(
|
|||||||
|
|
||||||
if processed_action.set_category:
|
if processed_action.set_category:
|
||||||
value = simple.eval(processed_action.set_category)
|
value = simple.eval(processed_action.set_category)
|
||||||
if isinstance(value, int):
|
if value is None:
|
||||||
|
transaction.category = None
|
||||||
|
elif isinstance(value, int):
|
||||||
transaction.category = TransactionCategory.objects.get(id=value)
|
transaction.category = TransactionCategory.objects.get(id=value)
|
||||||
else:
|
else:
|
||||||
transaction.category = TransactionCategory.objects.get(name=value)
|
transaction.category = TransactionCategory.objects.get(name=value)
|
||||||
@@ -458,7 +460,9 @@ def check_for_transaction_rules(
|
|||||||
transaction.account = account
|
transaction.account = account
|
||||||
|
|
||||||
elif field == TransactionRuleAction.Field.category:
|
elif field == TransactionRuleAction.Field.category:
|
||||||
if isinstance(new_value, int):
|
if new_value is None:
|
||||||
|
transaction.category = None
|
||||||
|
elif isinstance(new_value, int):
|
||||||
category = TransactionCategory.objects.get(id=new_value)
|
category = TransactionCategory.objects.get(id=new_value)
|
||||||
transaction.category = category
|
transaction.category = category
|
||||||
elif isinstance(new_value, str):
|
elif isinstance(new_value, str):
|
||||||
|
|||||||
@@ -0,0 +1 @@
|
|||||||
|
|
||||||
@@ -0,0 +1,82 @@
|
|||||||
|
from datetime import date
|
||||||
|
from decimal import Decimal
|
||||||
|
from unittest.mock import patch
|
||||||
|
|
||||||
|
from django.contrib.auth import get_user_model
|
||||||
|
from django.test import TransactionTestCase
|
||||||
|
|
||||||
|
from apps.accounts.models import Account
|
||||||
|
from apps.currencies.models import Currency
|
||||||
|
from apps.rules.models import TransactionRule, UpdateOrCreateTransactionRuleAction
|
||||||
|
from apps.rules.tasks import check_for_transaction_rules
|
||||||
|
from apps.transactions.models import Transaction
|
||||||
|
|
||||||
|
|
||||||
|
def run_check_for_transaction_rules_without_worker_wrapper(**kwargs):
|
||||||
|
task_func = check_for_transaction_rules.func
|
||||||
|
task_func = getattr(task_func, "__wrapped__", task_func)
|
||||||
|
|
||||||
|
return task_func(**kwargs)
|
||||||
|
|
||||||
|
|
||||||
|
class CheckForTransactionRulesTests(TransactionTestCase):
|
||||||
|
def setUp(self):
|
||||||
|
User = get_user_model()
|
||||||
|
self.user = User.objects.create_user(
|
||||||
|
email="rules@example.com",
|
||||||
|
password="testpass123",
|
||||||
|
)
|
||||||
|
self.currency = Currency.objects.create(
|
||||||
|
code="USD",
|
||||||
|
name="US Dollar",
|
||||||
|
decimal_places=2,
|
||||||
|
)
|
||||||
|
self.account = Account.objects.create(
|
||||||
|
name="Main Account",
|
||||||
|
currency=self.currency,
|
||||||
|
owner=self.user,
|
||||||
|
)
|
||||||
|
|
||||||
|
@patch("apps.rules.signals.check_for_transaction_rules.defer")
|
||||||
|
def test_update_or_create_action_can_clear_category_from_none_expression(
|
||||||
|
self, mock_defer
|
||||||
|
):
|
||||||
|
source_transaction = Transaction.objects.create(
|
||||||
|
account=self.account,
|
||||||
|
type=Transaction.Type.EXPENSE,
|
||||||
|
amount=Decimal("10.00"),
|
||||||
|
date=date(2026, 5, 4),
|
||||||
|
reference_date=date(2026, 5, 1),
|
||||||
|
description="Source without category",
|
||||||
|
category=None,
|
||||||
|
owner=self.user,
|
||||||
|
)
|
||||||
|
rule = TransactionRule.objects.create(
|
||||||
|
active=True,
|
||||||
|
on_create=False,
|
||||||
|
on_update=True,
|
||||||
|
name="Copy transaction",
|
||||||
|
trigger="True",
|
||||||
|
owner=self.user,
|
||||||
|
)
|
||||||
|
UpdateOrCreateTransactionRuleAction.objects.create(
|
||||||
|
rule=rule,
|
||||||
|
set_account="account_id",
|
||||||
|
set_type="'EX'",
|
||||||
|
set_date="date",
|
||||||
|
set_reference_date="reference_date",
|
||||||
|
set_amount="amount",
|
||||||
|
set_description="'Generated transaction'",
|
||||||
|
set_category="category_name",
|
||||||
|
)
|
||||||
|
|
||||||
|
run_check_for_transaction_rules_without_worker_wrapper(
|
||||||
|
instance_id=source_transaction.id,
|
||||||
|
user_id=self.user.id,
|
||||||
|
signal="transaction_updated",
|
||||||
|
)
|
||||||
|
|
||||||
|
generated_transaction = Transaction.objects.get(
|
||||||
|
description="Generated transaction"
|
||||||
|
)
|
||||||
|
self.assertIsNone(generated_transaction.category)
|
||||||
@@ -0,0 +1,473 @@
|
|||||||
|
"""Object-level authorization tests for the transaction-rule endpoints.
|
||||||
|
|
||||||
|
Regression coverage for GHSA-83g9-vjqf-2j5q: SharedObjectManager scopes rules to
|
||||||
|
what a user may *see*, which included other people's public and shared-with-them
|
||||||
|
rules. Several mutating endpoints treated that visibility as permission to write.
|
||||||
|
"""
|
||||||
|
|
||||||
|
from django.contrib.auth import get_user_model
|
||||||
|
from django.test import TestCase, override_settings
|
||||||
|
from django.urls import reverse
|
||||||
|
|
||||||
|
from apps.rules.models import (
|
||||||
|
TransactionRule,
|
||||||
|
TransactionRuleAction,
|
||||||
|
UpdateOrCreateTransactionRuleAction,
|
||||||
|
)
|
||||||
|
|
||||||
|
HTMX = {"HTTP_HX_REQUEST": "true"}
|
||||||
|
|
||||||
|
|
||||||
|
@override_settings(
|
||||||
|
STORAGES={
|
||||||
|
"default": {"BACKEND": "django.core.files.storage.FileSystemStorage"},
|
||||||
|
"staticfiles": {
|
||||||
|
"BACKEND": "django.contrib.staticfiles.storage.StaticFilesStorage"
|
||||||
|
},
|
||||||
|
},
|
||||||
|
WHITENOISE_AUTOREFRESH=True,
|
||||||
|
DEMO=False,
|
||||||
|
)
|
||||||
|
class TransactionRuleObjectPermissionTests(TestCase):
|
||||||
|
def setUp(self):
|
||||||
|
User = get_user_model()
|
||||||
|
self.owner = User.objects.create_user(
|
||||||
|
email="owner@test.com", password="testpass123"
|
||||||
|
)
|
||||||
|
self.shared_user = User.objects.create_user(
|
||||||
|
email="shared@test.com", password="testpass123"
|
||||||
|
)
|
||||||
|
self.stranger = User.objects.create_user(
|
||||||
|
email="stranger@test.com", password="testpass123"
|
||||||
|
)
|
||||||
|
|
||||||
|
# Public: visible to everyone through SharedObjectManager.
|
||||||
|
self.public_rule = TransactionRule.all_objects.create(
|
||||||
|
name="Public rule",
|
||||||
|
trigger="True",
|
||||||
|
owner=self.owner,
|
||||||
|
visibility="public",
|
||||||
|
active=True,
|
||||||
|
)
|
||||||
|
# Private but shared with shared_user: visible to them, not to stranger.
|
||||||
|
self.shared_rule = TransactionRule.all_objects.create(
|
||||||
|
name="Shared rule",
|
||||||
|
trigger="True",
|
||||||
|
owner=self.owner,
|
||||||
|
visibility="private",
|
||||||
|
active=True,
|
||||||
|
)
|
||||||
|
self.shared_rule.shared_with.add(self.shared_user)
|
||||||
|
# Private and unshared: invisible to everyone but the owner.
|
||||||
|
self.private_rule = TransactionRule.all_objects.create(
|
||||||
|
name="Private rule",
|
||||||
|
trigger="True",
|
||||||
|
owner=self.owner,
|
||||||
|
visibility="private",
|
||||||
|
active=True,
|
||||||
|
)
|
||||||
|
|
||||||
|
self.public_action = TransactionRuleAction.objects.create(
|
||||||
|
rule=self.public_rule, field="notes", value="owned by owner"
|
||||||
|
)
|
||||||
|
self.private_action = TransactionRuleAction.objects.create(
|
||||||
|
rule=self.private_rule, field="notes", value="owned by owner"
|
||||||
|
)
|
||||||
|
self.public_linked_action = UpdateOrCreateTransactionRuleAction.objects.create(
|
||||||
|
rule=self.public_rule
|
||||||
|
)
|
||||||
|
self.private_linked_action = UpdateOrCreateTransactionRuleAction.objects.create(
|
||||||
|
rule=self.private_rule
|
||||||
|
)
|
||||||
|
|
||||||
|
def login(self, user):
|
||||||
|
self.client.force_login(user)
|
||||||
|
|
||||||
|
# ------------------------------------------------------------------
|
||||||
|
# toggle-active
|
||||||
|
# ------------------------------------------------------------------
|
||||||
|
def test_stranger_cannot_toggle_public_rule(self):
|
||||||
|
self.login(self.stranger)
|
||||||
|
|
||||||
|
response = self.client.get(
|
||||||
|
reverse(
|
||||||
|
"transaction_rule_toggle_activity",
|
||||||
|
kwargs={"transaction_rule_id": self.public_rule.id},
|
||||||
|
),
|
||||||
|
**HTMX,
|
||||||
|
)
|
||||||
|
|
||||||
|
self.assertEqual(response.status_code, 403)
|
||||||
|
self.public_rule.refresh_from_db()
|
||||||
|
self.assertTrue(self.public_rule.active)
|
||||||
|
|
||||||
|
def test_shared_user_cannot_toggle_shared_rule(self):
|
||||||
|
self.login(self.shared_user)
|
||||||
|
|
||||||
|
response = self.client.get(
|
||||||
|
reverse(
|
||||||
|
"transaction_rule_toggle_activity",
|
||||||
|
kwargs={"transaction_rule_id": self.shared_rule.id},
|
||||||
|
),
|
||||||
|
**HTMX,
|
||||||
|
)
|
||||||
|
|
||||||
|
self.assertEqual(response.status_code, 403)
|
||||||
|
self.shared_rule.refresh_from_db()
|
||||||
|
self.assertTrue(self.shared_rule.active)
|
||||||
|
|
||||||
|
def test_owner_can_toggle_own_rule(self):
|
||||||
|
self.login(self.owner)
|
||||||
|
|
||||||
|
response = self.client.get(
|
||||||
|
reverse(
|
||||||
|
"transaction_rule_toggle_activity",
|
||||||
|
kwargs={"transaction_rule_id": self.public_rule.id},
|
||||||
|
),
|
||||||
|
**HTMX,
|
||||||
|
)
|
||||||
|
|
||||||
|
self.assertEqual(response.status_code, 204)
|
||||||
|
self.public_rule.refresh_from_db()
|
||||||
|
self.assertFalse(self.public_rule.active)
|
||||||
|
|
||||||
|
# ------------------------------------------------------------------
|
||||||
|
# rule actions
|
||||||
|
# ------------------------------------------------------------------
|
||||||
|
def test_stranger_cannot_add_action_to_public_rule(self):
|
||||||
|
self.login(self.stranger)
|
||||||
|
|
||||||
|
response = self.client.post(
|
||||||
|
reverse(
|
||||||
|
"transaction_rule_action_add",
|
||||||
|
kwargs={"transaction_rule_id": self.public_rule.id},
|
||||||
|
),
|
||||||
|
data={"field": "notes", "value": "injected", "order": 0},
|
||||||
|
**HTMX,
|
||||||
|
)
|
||||||
|
|
||||||
|
self.assertEqual(response.status_code, 403)
|
||||||
|
self.assertFalse(
|
||||||
|
TransactionRuleAction.objects.filter(value="injected").exists()
|
||||||
|
)
|
||||||
|
|
||||||
|
def test_stranger_cannot_edit_action_on_public_rule(self):
|
||||||
|
self.login(self.stranger)
|
||||||
|
|
||||||
|
response = self.client.post(
|
||||||
|
reverse(
|
||||||
|
"transaction_rule_action_edit",
|
||||||
|
kwargs={"transaction_rule_action_id": self.public_action.id},
|
||||||
|
),
|
||||||
|
data={"field": "notes", "value": "rewritten", "order": 0},
|
||||||
|
**HTMX,
|
||||||
|
)
|
||||||
|
|
||||||
|
self.assertEqual(response.status_code, 403)
|
||||||
|
self.public_action.refresh_from_db()
|
||||||
|
self.assertEqual(self.public_action.value, "owned by owner")
|
||||||
|
|
||||||
|
def test_stranger_cannot_delete_action_on_public_rule(self):
|
||||||
|
self.login(self.stranger)
|
||||||
|
|
||||||
|
response = self.client.delete(
|
||||||
|
reverse(
|
||||||
|
"transaction_rule_action_delete",
|
||||||
|
kwargs={"transaction_rule_action_id": self.public_action.id},
|
||||||
|
),
|
||||||
|
**HTMX,
|
||||||
|
)
|
||||||
|
|
||||||
|
self.assertEqual(response.status_code, 403)
|
||||||
|
self.assertTrue(
|
||||||
|
TransactionRuleAction.objects.filter(pk=self.public_action.pk).exists()
|
||||||
|
)
|
||||||
|
|
||||||
|
def test_stranger_cannot_delete_action_on_invisible_rule(self):
|
||||||
|
"""The child manager is unscoped, so the id is reachable by guessing.
|
||||||
|
|
||||||
|
The parent rule is invisible to the stranger, so the response must be a
|
||||||
|
404 rather than a 403 that confirms the action exists.
|
||||||
|
"""
|
||||||
|
self.login(self.stranger)
|
||||||
|
|
||||||
|
response = self.client.delete(
|
||||||
|
reverse(
|
||||||
|
"transaction_rule_action_delete",
|
||||||
|
kwargs={"transaction_rule_action_id": self.private_action.id},
|
||||||
|
),
|
||||||
|
**HTMX,
|
||||||
|
)
|
||||||
|
|
||||||
|
self.assertEqual(response.status_code, 404)
|
||||||
|
self.assertTrue(
|
||||||
|
TransactionRuleAction.objects.filter(pk=self.private_action.pk).exists()
|
||||||
|
)
|
||||||
|
|
||||||
|
def test_owner_can_delete_own_action(self):
|
||||||
|
self.login(self.owner)
|
||||||
|
|
||||||
|
response = self.client.delete(
|
||||||
|
reverse(
|
||||||
|
"transaction_rule_action_delete",
|
||||||
|
kwargs={"transaction_rule_action_id": self.public_action.id},
|
||||||
|
),
|
||||||
|
**HTMX,
|
||||||
|
)
|
||||||
|
|
||||||
|
self.assertEqual(response.status_code, 204)
|
||||||
|
self.assertFalse(
|
||||||
|
TransactionRuleAction.objects.filter(pk=self.public_action.pk).exists()
|
||||||
|
)
|
||||||
|
|
||||||
|
# ------------------------------------------------------------------
|
||||||
|
# update-or-create rule actions
|
||||||
|
# ------------------------------------------------------------------
|
||||||
|
def test_stranger_cannot_add_linked_action_to_public_rule(self):
|
||||||
|
self.login(self.stranger)
|
||||||
|
|
||||||
|
response = self.client.post(
|
||||||
|
reverse(
|
||||||
|
"update_or_create_transaction_rule_action_add",
|
||||||
|
kwargs={"transaction_rule_id": self.public_rule.id},
|
||||||
|
),
|
||||||
|
data={},
|
||||||
|
**HTMX,
|
||||||
|
)
|
||||||
|
|
||||||
|
self.assertEqual(response.status_code, 403)
|
||||||
|
self.assertEqual(
|
||||||
|
UpdateOrCreateTransactionRuleAction.objects.filter(
|
||||||
|
rule=self.public_rule
|
||||||
|
).count(),
|
||||||
|
1,
|
||||||
|
)
|
||||||
|
|
||||||
|
def test_stranger_cannot_edit_linked_action_on_public_rule(self):
|
||||||
|
self.login(self.stranger)
|
||||||
|
|
||||||
|
response = self.client.post(
|
||||||
|
reverse(
|
||||||
|
"update_or_create_transaction_rule_action_edit",
|
||||||
|
kwargs={"pk": self.public_linked_action.id},
|
||||||
|
),
|
||||||
|
data={"filter": "injected"},
|
||||||
|
**HTMX,
|
||||||
|
)
|
||||||
|
|
||||||
|
self.assertEqual(response.status_code, 403)
|
||||||
|
self.public_linked_action.refresh_from_db()
|
||||||
|
self.assertEqual(self.public_linked_action.filter, "")
|
||||||
|
|
||||||
|
def test_stranger_cannot_delete_linked_action_on_public_rule(self):
|
||||||
|
self.login(self.stranger)
|
||||||
|
|
||||||
|
response = self.client.delete(
|
||||||
|
reverse(
|
||||||
|
"update_or_create_transaction_rule_action_delete",
|
||||||
|
kwargs={"pk": self.public_linked_action.id},
|
||||||
|
),
|
||||||
|
**HTMX,
|
||||||
|
)
|
||||||
|
|
||||||
|
self.assertEqual(response.status_code, 403)
|
||||||
|
self.assertTrue(
|
||||||
|
UpdateOrCreateTransactionRuleAction.objects.filter(
|
||||||
|
pk=self.public_linked_action.pk
|
||||||
|
).exists()
|
||||||
|
)
|
||||||
|
|
||||||
|
def test_stranger_cannot_delete_linked_action_on_invisible_rule(self):
|
||||||
|
self.login(self.stranger)
|
||||||
|
|
||||||
|
response = self.client.delete(
|
||||||
|
reverse(
|
||||||
|
"update_or_create_transaction_rule_action_delete",
|
||||||
|
kwargs={"pk": self.private_linked_action.id},
|
||||||
|
),
|
||||||
|
**HTMX,
|
||||||
|
)
|
||||||
|
|
||||||
|
self.assertEqual(response.status_code, 404)
|
||||||
|
self.assertTrue(
|
||||||
|
UpdateOrCreateTransactionRuleAction.objects.filter(
|
||||||
|
pk=self.private_linked_action.pk
|
||||||
|
).exists()
|
||||||
|
)
|
||||||
|
|
||||||
|
# ------------------------------------------------------------------
|
||||||
|
# edit / share / dry-run
|
||||||
|
# ------------------------------------------------------------------
|
||||||
|
def test_stranger_cannot_edit_public_rule(self):
|
||||||
|
self.login(self.stranger)
|
||||||
|
|
||||||
|
response = self.client.post(
|
||||||
|
reverse(
|
||||||
|
"transaction_rule_edit",
|
||||||
|
kwargs={"transaction_rule_id": self.public_rule.id},
|
||||||
|
),
|
||||||
|
data={"name": "hijacked", "trigger": "True", "order": 0},
|
||||||
|
**HTMX,
|
||||||
|
)
|
||||||
|
|
||||||
|
self.assertEqual(response.status_code, 403)
|
||||||
|
self.public_rule.refresh_from_db()
|
||||||
|
self.assertEqual(self.public_rule.name, "Public rule")
|
||||||
|
|
||||||
|
def test_stranger_cannot_change_sharing_of_public_rule(self):
|
||||||
|
self.login(self.stranger)
|
||||||
|
|
||||||
|
response = self.client.post(
|
||||||
|
reverse(
|
||||||
|
"transaction_rule_share_settings", kwargs={"pk": self.public_rule.id}
|
||||||
|
),
|
||||||
|
data={"visibility": "private"},
|
||||||
|
**HTMX,
|
||||||
|
)
|
||||||
|
|
||||||
|
self.assertEqual(response.status_code, 403)
|
||||||
|
self.public_rule.refresh_from_db()
|
||||||
|
self.assertEqual(self.public_rule.visibility, "public")
|
||||||
|
|
||||||
|
def test_stranger_cannot_dry_run_public_rule(self):
|
||||||
|
self.login(self.stranger)
|
||||||
|
|
||||||
|
response = self.client.get(
|
||||||
|
reverse(
|
||||||
|
"transaction_rule_dry_run_created", kwargs={"pk": self.public_rule.id}
|
||||||
|
),
|
||||||
|
**HTMX,
|
||||||
|
)
|
||||||
|
|
||||||
|
self.assertEqual(response.status_code, 403)
|
||||||
|
|
||||||
|
# ------------------------------------------------------------------
|
||||||
|
# delete: sharing may be revoked, ownership may not be overridden
|
||||||
|
# ------------------------------------------------------------------
|
||||||
|
def test_stranger_cannot_delete_public_rule(self):
|
||||||
|
"""The original condition fell through to delete() for public rules."""
|
||||||
|
self.login(self.stranger)
|
||||||
|
|
||||||
|
response = self.client.delete(
|
||||||
|
reverse(
|
||||||
|
"transaction_rule_delete",
|
||||||
|
kwargs={"transaction_rule_id": self.public_rule.id},
|
||||||
|
),
|
||||||
|
**HTMX,
|
||||||
|
)
|
||||||
|
|
||||||
|
self.assertEqual(response.status_code, 403)
|
||||||
|
self.assertTrue(
|
||||||
|
TransactionRule.all_objects.filter(pk=self.public_rule.pk).exists()
|
||||||
|
)
|
||||||
|
|
||||||
|
def test_shared_user_deleting_only_revokes_their_own_access(self):
|
||||||
|
self.login(self.shared_user)
|
||||||
|
|
||||||
|
response = self.client.delete(
|
||||||
|
reverse(
|
||||||
|
"transaction_rule_delete",
|
||||||
|
kwargs={"transaction_rule_id": self.shared_rule.id},
|
||||||
|
),
|
||||||
|
**HTMX,
|
||||||
|
)
|
||||||
|
|
||||||
|
self.assertEqual(response.status_code, 204)
|
||||||
|
self.assertTrue(
|
||||||
|
TransactionRule.all_objects.filter(pk=self.shared_rule.pk).exists()
|
||||||
|
)
|
||||||
|
self.assertNotIn(self.shared_user, self.shared_rule.shared_with.all())
|
||||||
|
|
||||||
|
def test_owner_can_delete_own_rule(self):
|
||||||
|
self.login(self.owner)
|
||||||
|
|
||||||
|
response = self.client.delete(
|
||||||
|
reverse(
|
||||||
|
"transaction_rule_delete",
|
||||||
|
kwargs={"transaction_rule_id": self.public_rule.id},
|
||||||
|
),
|
||||||
|
**HTMX,
|
||||||
|
)
|
||||||
|
|
||||||
|
self.assertEqual(response.status_code, 204)
|
||||||
|
self.assertFalse(
|
||||||
|
TransactionRule.all_objects.filter(pk=self.public_rule.pk).exists()
|
||||||
|
)
|
||||||
|
|
||||||
|
# ------------------------------------------------------------------
|
||||||
|
# reads must keep working for shared users
|
||||||
|
# ------------------------------------------------------------------
|
||||||
|
def test_shared_user_can_still_view_shared_rule(self):
|
||||||
|
self.login(self.shared_user)
|
||||||
|
|
||||||
|
response = self.client.get(
|
||||||
|
reverse(
|
||||||
|
"transaction_rule_view",
|
||||||
|
kwargs={"transaction_rule_id": self.shared_rule.id},
|
||||||
|
),
|
||||||
|
**HTMX,
|
||||||
|
)
|
||||||
|
|
||||||
|
self.assertEqual(response.status_code, 200)
|
||||||
|
|
||||||
|
def test_stranger_can_view_public_rule(self):
|
||||||
|
self.login(self.stranger)
|
||||||
|
|
||||||
|
response = self.client.get(
|
||||||
|
reverse(
|
||||||
|
"transaction_rule_view",
|
||||||
|
kwargs={"transaction_rule_id": self.public_rule.id},
|
||||||
|
),
|
||||||
|
**HTMX,
|
||||||
|
)
|
||||||
|
|
||||||
|
self.assertEqual(response.status_code, 200)
|
||||||
|
|
||||||
|
def test_stranger_cannot_view_private_rule(self):
|
||||||
|
self.login(self.stranger)
|
||||||
|
|
||||||
|
response = self.client.get(
|
||||||
|
reverse(
|
||||||
|
"transaction_rule_view",
|
||||||
|
kwargs={"transaction_rule_id": self.private_rule.id},
|
||||||
|
),
|
||||||
|
**HTMX,
|
||||||
|
)
|
||||||
|
|
||||||
|
self.assertEqual(response.status_code, 404)
|
||||||
|
|
||||||
|
# ------------------------------------------------------------------
|
||||||
|
# take ownership stays available for unowned rules only
|
||||||
|
# ------------------------------------------------------------------
|
||||||
|
def test_take_ownership_of_unowned_rule_still_works(self):
|
||||||
|
unowned = TransactionRule.all_objects.create(
|
||||||
|
name="Legacy rule", trigger="True", owner=None, visibility="private"
|
||||||
|
)
|
||||||
|
self.login(self.stranger)
|
||||||
|
|
||||||
|
response = self.client.get(
|
||||||
|
reverse(
|
||||||
|
"transaction_rule_take_ownership",
|
||||||
|
kwargs={"transaction_rule_id": unowned.id},
|
||||||
|
),
|
||||||
|
**HTMX,
|
||||||
|
)
|
||||||
|
|
||||||
|
self.assertEqual(response.status_code, 204)
|
||||||
|
unowned.refresh_from_db()
|
||||||
|
self.assertEqual(unowned.owner, self.stranger)
|
||||||
|
|
||||||
|
def test_cannot_take_ownership_of_someone_elses_public_rule(self):
|
||||||
|
self.login(self.stranger)
|
||||||
|
|
||||||
|
response = self.client.get(
|
||||||
|
reverse(
|
||||||
|
"transaction_rule_take_ownership",
|
||||||
|
kwargs={"transaction_rule_id": self.public_rule.id},
|
||||||
|
),
|
||||||
|
**HTMX,
|
||||||
|
)
|
||||||
|
|
||||||
|
self.assertEqual(response.status_code, 403)
|
||||||
|
self.public_rule.refresh_from_db()
|
||||||
|
self.assertEqual(self.public_rule.owner, self.owner)
|
||||||
+58
-47
@@ -4,13 +4,19 @@ from copy import deepcopy
|
|||||||
|
|
||||||
from django.contrib import messages
|
from django.contrib import messages
|
||||||
from django.contrib.auth.decorators import login_required
|
from django.contrib.auth.decorators import login_required
|
||||||
|
from django.core.exceptions import PermissionDenied
|
||||||
from django.db import transaction
|
from django.db import transaction
|
||||||
from django.http import HttpResponse
|
from django.http import HttpResponse
|
||||||
from django.shortcuts import render, get_object_or_404, redirect
|
from django.shortcuts import render, redirect
|
||||||
from django.utils.translation import gettext_lazy as _
|
from django.utils.translation import gettext_lazy as _
|
||||||
from django.views.decorators.http import require_http_methods
|
from django.views.decorators.http import require_http_methods
|
||||||
|
|
||||||
from apps.common.decorators.htmx import only_htmx
|
from apps.common.decorators.htmx import only_htmx
|
||||||
|
from apps.common.functions.permissions import (
|
||||||
|
EDIT,
|
||||||
|
READ,
|
||||||
|
get_shared_object_or_error,
|
||||||
|
)
|
||||||
from apps.rules.forms import (
|
from apps.rules.forms import (
|
||||||
TransactionRuleForm,
|
TransactionRuleForm,
|
||||||
TransactionRuleActionForm,
|
TransactionRuleActionForm,
|
||||||
@@ -62,7 +68,9 @@ def rules_list(request):
|
|||||||
@disabled_on_demo
|
@disabled_on_demo
|
||||||
@require_http_methods(["GET", "POST"])
|
@require_http_methods(["GET", "POST"])
|
||||||
def transaction_rule_toggle_activity(request, transaction_rule_id, **kwargs):
|
def transaction_rule_toggle_activity(request, transaction_rule_id, **kwargs):
|
||||||
transaction_rule = get_object_or_404(TransactionRule, id=transaction_rule_id)
|
transaction_rule = get_shared_object_or_error(
|
||||||
|
TransactionRule, request, id=transaction_rule_id, level=EDIT
|
||||||
|
)
|
||||||
current_active = transaction_rule.active
|
current_active = transaction_rule.active
|
||||||
transaction_rule.active = not current_active
|
transaction_rule.active = not current_active
|
||||||
transaction_rule.save(update_fields=["active"])
|
transaction_rule.save(update_fields=["active"])
|
||||||
@@ -112,17 +120,9 @@ def transaction_rule_add(request, **kwargs):
|
|||||||
@disabled_on_demo
|
@disabled_on_demo
|
||||||
@require_http_methods(["GET", "POST"])
|
@require_http_methods(["GET", "POST"])
|
||||||
def transaction_rule_edit(request, transaction_rule_id):
|
def transaction_rule_edit(request, transaction_rule_id):
|
||||||
transaction_rule = get_object_or_404(TransactionRule, id=transaction_rule_id)
|
transaction_rule = get_shared_object_or_error(
|
||||||
|
TransactionRule, request, id=transaction_rule_id, level=EDIT
|
||||||
if transaction_rule.owner and transaction_rule.owner != request.user:
|
)
|
||||||
messages.error(request, _("Only the owner can edit this"))
|
|
||||||
|
|
||||||
return HttpResponse(
|
|
||||||
status=204,
|
|
||||||
headers={
|
|
||||||
"HX-Trigger": "updated, hide_offcanvas",
|
|
||||||
},
|
|
||||||
)
|
|
||||||
|
|
||||||
if request.method == "POST":
|
if request.method == "POST":
|
||||||
form = TransactionRuleForm(request.POST, instance=transaction_rule)
|
form = TransactionRuleForm(request.POST, instance=transaction_rule)
|
||||||
@@ -151,7 +151,9 @@ def transaction_rule_edit(request, transaction_rule_id):
|
|||||||
@disabled_on_demo
|
@disabled_on_demo
|
||||||
@require_http_methods(["GET", "POST"])
|
@require_http_methods(["GET", "POST"])
|
||||||
def transaction_rule_view(request, transaction_rule_id):
|
def transaction_rule_view(request, transaction_rule_id):
|
||||||
transaction_rule = get_object_or_404(TransactionRule, id=transaction_rule_id)
|
transaction_rule = get_shared_object_or_error(
|
||||||
|
TransactionRule, request, id=transaction_rule_id, level=READ
|
||||||
|
)
|
||||||
|
|
||||||
edit_actions = transaction_rule.transaction_actions.all()
|
edit_actions = transaction_rule.transaction_actions.all()
|
||||||
update_or_create_actions = (
|
update_or_create_actions = (
|
||||||
@@ -175,17 +177,20 @@ def transaction_rule_view(request, transaction_rule_id):
|
|||||||
@disabled_on_demo
|
@disabled_on_demo
|
||||||
@require_http_methods(["DELETE"])
|
@require_http_methods(["DELETE"])
|
||||||
def transaction_rule_delete(request, transaction_rule_id):
|
def transaction_rule_delete(request, transaction_rule_id):
|
||||||
transaction_rule = get_object_or_404(TransactionRule, id=transaction_rule_id)
|
transaction_rule = get_shared_object_or_error(
|
||||||
|
TransactionRule, request, id=transaction_rule_id, level=READ
|
||||||
|
)
|
||||||
|
|
||||||
if (
|
if transaction_rule.is_editable_by(request.user):
|
||||||
transaction_rule.owner != request.user
|
transaction_rule.delete()
|
||||||
and request.user in transaction_rule.shared_with.all()
|
messages.success(request, _("Rule deleted successfully"))
|
||||||
):
|
elif transaction_rule.shared_with.filter(pk=request.user.pk).exists():
|
||||||
|
# Someone else's rule shared with us: we can drop our own access to it,
|
||||||
|
# but never delete it.
|
||||||
transaction_rule.shared_with.remove(request.user)
|
transaction_rule.shared_with.remove(request.user)
|
||||||
messages.success(request, _("Item no longer shared with you"))
|
messages.success(request, _("Item no longer shared with you"))
|
||||||
else:
|
else:
|
||||||
transaction_rule.delete()
|
raise PermissionDenied
|
||||||
messages.success(request, _("Rule deleted successfully"))
|
|
||||||
|
|
||||||
return HttpResponse(
|
return HttpResponse(
|
||||||
status=204,
|
status=204,
|
||||||
@@ -200,7 +205,9 @@ def transaction_rule_delete(request, transaction_rule_id):
|
|||||||
@disabled_on_demo
|
@disabled_on_demo
|
||||||
@require_http_methods(["GET"])
|
@require_http_methods(["GET"])
|
||||||
def transaction_rule_take_ownership(request, transaction_rule_id):
|
def transaction_rule_take_ownership(request, transaction_rule_id):
|
||||||
transaction_rule = get_object_or_404(TransactionRule, id=transaction_rule_id)
|
transaction_rule = get_shared_object_or_error(
|
||||||
|
TransactionRule, request, id=transaction_rule_id, level=EDIT
|
||||||
|
)
|
||||||
|
|
||||||
if not transaction_rule.owner:
|
if not transaction_rule.owner:
|
||||||
transaction_rule.owner = request.user
|
transaction_rule.owner = request.user
|
||||||
@@ -222,17 +229,7 @@ def transaction_rule_take_ownership(request, transaction_rule_id):
|
|||||||
@disabled_on_demo
|
@disabled_on_demo
|
||||||
@require_http_methods(["GET", "POST"])
|
@require_http_methods(["GET", "POST"])
|
||||||
def transaction_rule_share(request, pk):
|
def transaction_rule_share(request, pk):
|
||||||
obj = get_object_or_404(TransactionRule, id=pk)
|
obj = get_shared_object_or_error(TransactionRule, request, id=pk, level=EDIT)
|
||||||
|
|
||||||
if obj.owner and obj.owner != request.user:
|
|
||||||
messages.error(request, _("Only the owner can edit this"))
|
|
||||||
|
|
||||||
return HttpResponse(
|
|
||||||
status=204,
|
|
||||||
headers={
|
|
||||||
"HX-Trigger": "updated, hide_offcanvas",
|
|
||||||
},
|
|
||||||
)
|
|
||||||
|
|
||||||
if request.method == "POST":
|
if request.method == "POST":
|
||||||
form = SharedObjectForm(request.POST, instance=obj, user=request.user)
|
form = SharedObjectForm(request.POST, instance=obj, user=request.user)
|
||||||
@@ -261,7 +258,9 @@ def transaction_rule_share(request, pk):
|
|||||||
@disabled_on_demo
|
@disabled_on_demo
|
||||||
@require_http_methods(["GET", "POST"])
|
@require_http_methods(["GET", "POST"])
|
||||||
def transaction_rule_action_add(request, transaction_rule_id):
|
def transaction_rule_action_add(request, transaction_rule_id):
|
||||||
transaction_rule = get_object_or_404(TransactionRule, id=transaction_rule_id)
|
transaction_rule = get_shared_object_or_error(
|
||||||
|
TransactionRule, request, id=transaction_rule_id, level=EDIT
|
||||||
|
)
|
||||||
|
|
||||||
if request.method == "POST":
|
if request.method == "POST":
|
||||||
form = TransactionRuleActionForm(request.POST, rule=transaction_rule)
|
form = TransactionRuleActionForm(request.POST, rule=transaction_rule)
|
||||||
@@ -289,12 +288,14 @@ def transaction_rule_action_add(request, transaction_rule_id):
|
|||||||
@disabled_on_demo
|
@disabled_on_demo
|
||||||
@require_http_methods(["GET", "POST"])
|
@require_http_methods(["GET", "POST"])
|
||||||
def transaction_rule_action_edit(request, transaction_rule_action_id):
|
def transaction_rule_action_edit(request, transaction_rule_action_id):
|
||||||
transaction_rule_action = get_object_or_404(
|
transaction_rule_action = get_shared_object_or_error(
|
||||||
TransactionRuleAction, id=transaction_rule_action_id
|
TransactionRuleAction,
|
||||||
)
|
request,
|
||||||
transaction_rule = get_object_or_404(
|
id=transaction_rule_action_id,
|
||||||
TransactionRule, id=transaction_rule_action.rule.id
|
level=EDIT,
|
||||||
|
via="rule",
|
||||||
)
|
)
|
||||||
|
transaction_rule = transaction_rule_action.rule
|
||||||
|
|
||||||
if request.method == "POST":
|
if request.method == "POST":
|
||||||
form = TransactionRuleActionForm(
|
form = TransactionRuleActionForm(
|
||||||
@@ -327,8 +328,12 @@ def transaction_rule_action_edit(request, transaction_rule_action_id):
|
|||||||
@disabled_on_demo
|
@disabled_on_demo
|
||||||
@require_http_methods(["DELETE"])
|
@require_http_methods(["DELETE"])
|
||||||
def transaction_rule_action_delete(request, transaction_rule_action_id):
|
def transaction_rule_action_delete(request, transaction_rule_action_id):
|
||||||
transaction_rule_action = get_object_or_404(
|
transaction_rule_action = get_shared_object_or_error(
|
||||||
TransactionRuleAction, id=transaction_rule_action_id
|
TransactionRuleAction,
|
||||||
|
request,
|
||||||
|
id=transaction_rule_action_id,
|
||||||
|
level=EDIT,
|
||||||
|
via="rule",
|
||||||
)
|
)
|
||||||
|
|
||||||
transaction_rule_action.delete()
|
transaction_rule_action.delete()
|
||||||
@@ -348,7 +353,9 @@ def transaction_rule_action_delete(request, transaction_rule_action_id):
|
|||||||
@disabled_on_demo
|
@disabled_on_demo
|
||||||
@require_http_methods(["GET", "POST"])
|
@require_http_methods(["GET", "POST"])
|
||||||
def update_or_create_transaction_rule_action_add(request, transaction_rule_id):
|
def update_or_create_transaction_rule_action_add(request, transaction_rule_id):
|
||||||
transaction_rule = get_object_or_404(TransactionRule, id=transaction_rule_id)
|
transaction_rule = get_shared_object_or_error(
|
||||||
|
TransactionRule, request, id=transaction_rule_id, level=EDIT
|
||||||
|
)
|
||||||
|
|
||||||
if request.method == "POST":
|
if request.method == "POST":
|
||||||
form = UpdateOrCreateTransactionRuleActionForm(
|
form = UpdateOrCreateTransactionRuleActionForm(
|
||||||
@@ -380,7 +387,9 @@ def update_or_create_transaction_rule_action_add(request, transaction_rule_id):
|
|||||||
@disabled_on_demo
|
@disabled_on_demo
|
||||||
@require_http_methods(["GET", "POST"])
|
@require_http_methods(["GET", "POST"])
|
||||||
def update_or_create_transaction_rule_action_edit(request, pk):
|
def update_or_create_transaction_rule_action_edit(request, pk):
|
||||||
linked_action = get_object_or_404(UpdateOrCreateTransactionRuleAction, id=pk)
|
linked_action = get_shared_object_or_error(
|
||||||
|
UpdateOrCreateTransactionRuleAction, request, id=pk, level=EDIT, via="rule"
|
||||||
|
)
|
||||||
transaction_rule = linked_action.rule
|
transaction_rule = linked_action.rule
|
||||||
|
|
||||||
if request.method == "POST":
|
if request.method == "POST":
|
||||||
@@ -415,7 +424,9 @@ def update_or_create_transaction_rule_action_edit(request, pk):
|
|||||||
@disabled_on_demo
|
@disabled_on_demo
|
||||||
@require_http_methods(["DELETE"])
|
@require_http_methods(["DELETE"])
|
||||||
def update_or_create_transaction_rule_action_delete(request, pk):
|
def update_or_create_transaction_rule_action_delete(request, pk):
|
||||||
linked_action = get_object_or_404(UpdateOrCreateTransactionRuleAction, id=pk)
|
linked_action = get_shared_object_or_error(
|
||||||
|
UpdateOrCreateTransactionRuleAction, request, id=pk, level=EDIT, via="rule"
|
||||||
|
)
|
||||||
|
|
||||||
linked_action.delete()
|
linked_action.delete()
|
||||||
|
|
||||||
@@ -436,7 +447,7 @@ def update_or_create_transaction_rule_action_delete(request, pk):
|
|||||||
@disabled_on_demo
|
@disabled_on_demo
|
||||||
@require_http_methods(["GET", "POST"])
|
@require_http_methods(["GET", "POST"])
|
||||||
def dry_run_rule_created(request, pk):
|
def dry_run_rule_created(request, pk):
|
||||||
rule = get_object_or_404(TransactionRule, id=pk)
|
rule = get_shared_object_or_error(TransactionRule, request, id=pk, level=EDIT)
|
||||||
logs = None
|
logs = None
|
||||||
results = None
|
results = None
|
||||||
|
|
||||||
@@ -481,7 +492,7 @@ def dry_run_rule_created(request, pk):
|
|||||||
@disabled_on_demo
|
@disabled_on_demo
|
||||||
@require_http_methods(["GET", "POST"])
|
@require_http_methods(["GET", "POST"])
|
||||||
def dry_run_rule_deleted(request, pk):
|
def dry_run_rule_deleted(request, pk):
|
||||||
rule = get_object_or_404(TransactionRule, id=pk)
|
rule = get_shared_object_or_error(TransactionRule, request, id=pk, level=EDIT)
|
||||||
logs = None
|
logs = None
|
||||||
results = None
|
results = None
|
||||||
|
|
||||||
@@ -526,7 +537,7 @@ def dry_run_rule_deleted(request, pk):
|
|||||||
@disabled_on_demo
|
@disabled_on_demo
|
||||||
@require_http_methods(["GET", "POST"])
|
@require_http_methods(["GET", "POST"])
|
||||||
def dry_run_rule_updated(request, pk):
|
def dry_run_rule_updated(request, pk):
|
||||||
rule = get_object_or_404(TransactionRule, id=pk)
|
rule = get_shared_object_or_error(TransactionRule, request, id=pk, level=EDIT)
|
||||||
logs = None
|
logs = None
|
||||||
results = None
|
results = None
|
||||||
|
|
||||||
|
|||||||
@@ -23,6 +23,11 @@ SITUACAO_CHOICES = (
|
|||||||
("0", _("Projected")),
|
("0", _("Projected")),
|
||||||
)
|
)
|
||||||
|
|
||||||
|
MUTE_STATUS_CHOICES = (
|
||||||
|
("active", _("Active")),
|
||||||
|
("muted", _("Muted")),
|
||||||
|
)
|
||||||
|
|
||||||
|
|
||||||
def content_filter(queryset, name, value):
|
def content_filter(queryset, name, value):
|
||||||
queryset = queryset.filter(
|
queryset = queryset.filter(
|
||||||
@@ -36,6 +41,12 @@ class MonthYearFilter(Filter):
|
|||||||
|
|
||||||
|
|
||||||
class TransactionsFilter(django_filters.FilterSet):
|
class TransactionsFilter(django_filters.FilterSet):
|
||||||
|
default_filter_values = {
|
||||||
|
"type": {"IN", "EX"},
|
||||||
|
"is_paid": {"1", "0"},
|
||||||
|
"mute_status": {"active", "muted"},
|
||||||
|
}
|
||||||
|
|
||||||
description = django_filters.CharFilter(
|
description = django_filters.CharFilter(
|
||||||
label=_("Content"),
|
label=_("Content"),
|
||||||
method=content_filter,
|
method=content_filter,
|
||||||
@@ -78,6 +89,11 @@ class TransactionsFilter(django_filters.FilterSet):
|
|||||||
choices=SITUACAO_CHOICES,
|
choices=SITUACAO_CHOICES,
|
||||||
field_name="is_paid",
|
field_name="is_paid",
|
||||||
)
|
)
|
||||||
|
mute_status = django_filters.MultipleChoiceFilter(
|
||||||
|
choices=MUTE_STATUS_CHOICES,
|
||||||
|
method="filter_mute_status",
|
||||||
|
label=_("Mute Status"),
|
||||||
|
)
|
||||||
date_start = django_filters.DateFilter(
|
date_start = django_filters.DateFilter(
|
||||||
field_name="date",
|
field_name="date",
|
||||||
lookup_expr="gte",
|
lookup_expr="gte",
|
||||||
@@ -140,6 +156,9 @@ class TransactionsFilter(django_filters.FilterSet):
|
|||||||
if data.get("is_paid") is None:
|
if data.get("is_paid") is None:
|
||||||
data.setlist("is_paid", ["1", "0"])
|
data.setlist("is_paid", ["1", "0"])
|
||||||
|
|
||||||
|
if data.get("mute_status") is None:
|
||||||
|
data.setlist("mute_status", ["active", "muted"])
|
||||||
|
|
||||||
super().__init__(data, *args, **kwargs)
|
super().__init__(data, *args, **kwargs)
|
||||||
|
|
||||||
self.form.helper = FormHelper()
|
self.form.helper = FormHelper()
|
||||||
@@ -155,6 +174,10 @@ class TransactionsFilter(django_filters.FilterSet):
|
|||||||
"is_paid",
|
"is_paid",
|
||||||
template="transactions/widgets/transaction_type_filter_buttons.html",
|
template="transactions/widgets/transaction_type_filter_buttons.html",
|
||||||
),
|
),
|
||||||
|
Field(
|
||||||
|
"mute_status",
|
||||||
|
template="transactions/widgets/transaction_type_filter_buttons.html",
|
||||||
|
),
|
||||||
Field("description"),
|
Field("description"),
|
||||||
Row(Column("date_start"), Column("date_end")),
|
Row(Column("date_start"), Column("date_end")),
|
||||||
Row(
|
Row(
|
||||||
@@ -200,6 +223,23 @@ class TransactionsFilter(django_filters.FilterSet):
|
|||||||
]
|
]
|
||||||
self.form.fields["entities"].choices = custom_entity_choices + entity_choices
|
self.form.fields["entities"].choices = custom_entity_choices + entity_choices
|
||||||
|
|
||||||
|
@property
|
||||||
|
def has_active_filters(self):
|
||||||
|
for name in self.base_filters:
|
||||||
|
if hasattr(self.form.data, "getlist"):
|
||||||
|
values = self.form.data.getlist(name)
|
||||||
|
else:
|
||||||
|
value = self.form.data.get(name)
|
||||||
|
values = value if isinstance(value, (list, tuple)) else [value]
|
||||||
|
|
||||||
|
values = [str(value) for value in values if value not in (None, "")]
|
||||||
|
if not values:
|
||||||
|
continue
|
||||||
|
if set(values) == self.default_filter_values.get(name):
|
||||||
|
continue
|
||||||
|
return True
|
||||||
|
return False
|
||||||
|
|
||||||
@staticmethod
|
@staticmethod
|
||||||
def filter_category(queryset, name, value):
|
def filter_category(queryset, name, value):
|
||||||
if not value:
|
if not value:
|
||||||
@@ -268,3 +308,36 @@ class TransactionsFilter(django_filters.FilterSet):
|
|||||||
return queryset.filter(q).distinct()
|
return queryset.filter(q).distinct()
|
||||||
|
|
||||||
return queryset
|
return queryset
|
||||||
|
|
||||||
|
@staticmethod
|
||||||
|
def filter_mute_status(queryset, name, value):
|
||||||
|
from apps.common.middleware.thread_local import get_current_user
|
||||||
|
|
||||||
|
if not value:
|
||||||
|
return queryset
|
||||||
|
|
||||||
|
value = list(value)
|
||||||
|
|
||||||
|
# If both are selected, return all
|
||||||
|
if "active" in value and "muted" in value:
|
||||||
|
return queryset
|
||||||
|
|
||||||
|
user = get_current_user()
|
||||||
|
|
||||||
|
# Only Active selected: exclude muted transactions
|
||||||
|
if "active" in value:
|
||||||
|
return (
|
||||||
|
queryset.exclude(account__untracked_by=user)
|
||||||
|
.filter(
|
||||||
|
mute=False,
|
||||||
|
)
|
||||||
|
.filter(Q(category__mute=False) | Q(category__isnull=True))
|
||||||
|
)
|
||||||
|
|
||||||
|
# Only Muted selected: include only muted transactions
|
||||||
|
if "muted" in value:
|
||||||
|
return queryset.filter(
|
||||||
|
Q(account__untracked_by=user) | Q(category__mute=True) | Q(mute=True)
|
||||||
|
)
|
||||||
|
|
||||||
|
return queryset
|
||||||
|
|||||||
@@ -5,6 +5,7 @@ from apps.common.fields.forms.dynamic_select import (
|
|||||||
DynamicModelChoiceField,
|
DynamicModelChoiceField,
|
||||||
DynamicModelMultipleChoiceField,
|
DynamicModelMultipleChoiceField,
|
||||||
)
|
)
|
||||||
|
from apps.common.middleware.thread_local import get_current_user
|
||||||
from apps.common.widgets.crispy.daisyui import Switch
|
from apps.common.widgets.crispy.daisyui import Switch
|
||||||
from apps.common.widgets.crispy.submit import NoClassSubmit
|
from apps.common.widgets.crispy.submit import NoClassSubmit
|
||||||
from apps.common.widgets.datepicker import AirDatePickerInput, AirMonthYearPickerInput
|
from apps.common.widgets.datepicker import AirDatePickerInput, AirMonthYearPickerInput
|
||||||
@@ -13,6 +14,7 @@ from apps.common.widgets.tom_select import TomSelect
|
|||||||
from apps.rules.signals import transaction_created, transaction_updated
|
from apps.rules.signals import transaction_created, transaction_updated
|
||||||
from apps.transactions.models import (
|
from apps.transactions.models import (
|
||||||
InstallmentPlan,
|
InstallmentPlan,
|
||||||
|
TransactionAttachment,
|
||||||
QuickTransaction,
|
QuickTransaction,
|
||||||
RecurringTransaction,
|
RecurringTransaction,
|
||||||
Transaction,
|
Transaction,
|
||||||
@@ -35,6 +37,22 @@ from django.db.models import Q
|
|||||||
from django.utils.translation import gettext_lazy as _
|
from django.utils.translation import gettext_lazy as _
|
||||||
|
|
||||||
|
|
||||||
|
class MultipleFileInput(forms.ClearableFileInput):
|
||||||
|
allow_multiple_selected = True
|
||||||
|
|
||||||
|
|
||||||
|
class MultipleFileField(forms.FileField):
|
||||||
|
widget = MultipleFileInput
|
||||||
|
|
||||||
|
def clean(self, data, initial=None):
|
||||||
|
single_file_clean = super().clean
|
||||||
|
if isinstance(data, (list, tuple)):
|
||||||
|
return [single_file_clean(file, initial) for file in data]
|
||||||
|
if data:
|
||||||
|
return [single_file_clean(data, initial)]
|
||||||
|
return []
|
||||||
|
|
||||||
|
|
||||||
class TransactionForm(forms.ModelForm):
|
class TransactionForm(forms.ModelForm):
|
||||||
category = DynamicModelChoiceField(
|
category = DynamicModelChoiceField(
|
||||||
create_field="name",
|
create_field="name",
|
||||||
@@ -116,6 +134,9 @@ class TransactionForm(forms.ModelForm):
|
|||||||
self.fields["account"].queryset = Account.objects.filter(
|
self.fields["account"].queryset = Account.objects.filter(
|
||||||
is_archived=False,
|
is_archived=False,
|
||||||
)
|
)
|
||||||
|
user_settings = get_current_user().settings
|
||||||
|
if user_settings.default_account:
|
||||||
|
self.fields["account"].initial = user_settings.default_account
|
||||||
|
|
||||||
self.fields["category"].queryset = TransactionCategory.objects.filter(
|
self.fields["category"].queryset = TransactionCategory.objects.filter(
|
||||||
active=True
|
active=True
|
||||||
@@ -243,6 +264,41 @@ class TransactionForm(forms.ModelForm):
|
|||||||
return instance
|
return instance
|
||||||
|
|
||||||
|
|
||||||
|
class TransactionAttachmentForm(forms.Form):
|
||||||
|
attachments = MultipleFileField(
|
||||||
|
required=True,
|
||||||
|
label=_("Attachments"),
|
||||||
|
help_text=_("Files are private and only visible to users with access to this transaction."),
|
||||||
|
)
|
||||||
|
|
||||||
|
def __init__(self, *args, **kwargs):
|
||||||
|
super().__init__(*args, **kwargs)
|
||||||
|
self.helper = FormHelper()
|
||||||
|
self.helper.form_tag = False
|
||||||
|
self.helper.form_method = "post"
|
||||||
|
self.helper.layout = Layout(
|
||||||
|
"attachments",
|
||||||
|
FormActions(
|
||||||
|
NoClassSubmit("submit", _("Upload"), css_class="btn btn-primary"),
|
||||||
|
),
|
||||||
|
)
|
||||||
|
|
||||||
|
def save(self, transaction, uploaded_by):
|
||||||
|
created = []
|
||||||
|
for attachment in self.cleaned_data.get("attachments") or []:
|
||||||
|
created.append(
|
||||||
|
TransactionAttachment.objects.create(
|
||||||
|
transaction=transaction,
|
||||||
|
file=attachment,
|
||||||
|
original_name=attachment.name,
|
||||||
|
content_type=getattr(attachment, "content_type", ""),
|
||||||
|
size=attachment.size,
|
||||||
|
uploaded_by=uploaded_by,
|
||||||
|
)
|
||||||
|
)
|
||||||
|
return created
|
||||||
|
|
||||||
|
|
||||||
class QuickTransactionForm(forms.ModelForm):
|
class QuickTransactionForm(forms.ModelForm):
|
||||||
category = DynamicModelChoiceField(
|
category = DynamicModelChoiceField(
|
||||||
create_field="name",
|
create_field="name",
|
||||||
@@ -768,6 +824,9 @@ class InstallmentPlanForm(forms.ModelForm):
|
|||||||
).distinct()
|
).distinct()
|
||||||
else:
|
else:
|
||||||
self.fields["account"].queryset = Account.objects.filter(is_archived=False)
|
self.fields["account"].queryset = Account.objects.filter(is_archived=False)
|
||||||
|
user_settings = get_current_user().settings
|
||||||
|
if user_settings.default_account:
|
||||||
|
self.fields["account"].initial = user_settings.default_account
|
||||||
|
|
||||||
self.fields["category"].queryset = TransactionCategory.objects.filter(
|
self.fields["category"].queryset = TransactionCategory.objects.filter(
|
||||||
active=True
|
active=True
|
||||||
@@ -1010,6 +1069,10 @@ class RecurringTransactionForm(forms.ModelForm):
|
|||||||
).distinct()
|
).distinct()
|
||||||
else:
|
else:
|
||||||
self.fields["account"].queryset = Account.objects.filter(is_archived=False)
|
self.fields["account"].queryset = Account.objects.filter(is_archived=False)
|
||||||
|
|
||||||
|
user_settings = get_current_user().settings
|
||||||
|
if user_settings.default_account:
|
||||||
|
self.fields["account"].initial = user_settings.default_account
|
||||||
|
|
||||||
self.fields["category"].queryset = TransactionCategory.objects.filter(
|
self.fields["category"].queryset = TransactionCategory.objects.filter(
|
||||||
active=True
|
active=True
|
||||||
|
|||||||
@@ -0,0 +1,38 @@
|
|||||||
|
# Generated by Django 5.2.13 on 2026-06-06 02:34
|
||||||
|
|
||||||
|
import apps.transactions.models
|
||||||
|
import apps.transactions.storage
|
||||||
|
import django.db.models.deletion
|
||||||
|
import uuid
|
||||||
|
from django.conf import settings
|
||||||
|
from django.db import migrations, models
|
||||||
|
|
||||||
|
|
||||||
|
class Migration(migrations.Migration):
|
||||||
|
|
||||||
|
dependencies = [
|
||||||
|
('transactions', '0048_recurringtransaction_keep_at_most'),
|
||||||
|
migrations.swappable_dependency(settings.AUTH_USER_MODEL),
|
||||||
|
]
|
||||||
|
|
||||||
|
operations = [
|
||||||
|
migrations.CreateModel(
|
||||||
|
name='TransactionAttachment',
|
||||||
|
fields=[
|
||||||
|
('id', models.UUIDField(default=uuid.uuid4, editable=False, primary_key=True, serialize=False)),
|
||||||
|
('file', models.FileField(storage=apps.transactions.storage.PrivateMediaStorage(), upload_to=apps.transactions.models.transaction_attachment_path, verbose_name='File')),
|
||||||
|
('original_name', models.CharField(max_length=255, verbose_name='Original Name')),
|
||||||
|
('content_type', models.CharField(blank=True, max_length=255, verbose_name='Content Type')),
|
||||||
|
('size', models.PositiveBigIntegerField(default=0, verbose_name='Size')),
|
||||||
|
('created_at', models.DateTimeField(auto_now_add=True)),
|
||||||
|
('transaction', models.ForeignKey(on_delete=django.db.models.deletion.CASCADE, related_name='attachments', to='transactions.transaction', verbose_name='Transaction')),
|
||||||
|
('uploaded_by', models.ForeignKey(on_delete=django.db.models.deletion.PROTECT, related_name='transaction_attachments', to=settings.AUTH_USER_MODEL, verbose_name='Uploaded By')),
|
||||||
|
],
|
||||||
|
options={
|
||||||
|
'verbose_name': 'Transaction Attachment',
|
||||||
|
'verbose_name_plural': 'Transaction Attachments',
|
||||||
|
'db_table': 'transaction_attachments',
|
||||||
|
'ordering': ['-created_at', 'original_name'],
|
||||||
|
},
|
||||||
|
),
|
||||||
|
]
|
||||||
@@ -0,0 +1,38 @@
|
|||||||
|
import django.db.models.deletion
|
||||||
|
from django.conf import settings
|
||||||
|
from django.db import migrations, models
|
||||||
|
|
||||||
|
|
||||||
|
class Migration(migrations.Migration):
|
||||||
|
dependencies = [
|
||||||
|
("transactions", "0049_transactionattachment"),
|
||||||
|
migrations.swappable_dependency(settings.AUTH_USER_MODEL),
|
||||||
|
]
|
||||||
|
|
||||||
|
operations = [
|
||||||
|
migrations.CreateModel(
|
||||||
|
name="FilterPreset",
|
||||||
|
fields=[
|
||||||
|
(
|
||||||
|
"id",
|
||||||
|
models.BigAutoField(
|
||||||
|
auto_created=True,
|
||||||
|
primary_key=True,
|
||||||
|
serialize=False,
|
||||||
|
verbose_name="ID",
|
||||||
|
),
|
||||||
|
),
|
||||||
|
("name", models.CharField(max_length=100)),
|
||||||
|
("parameters", models.JSONField(default=dict)),
|
||||||
|
(
|
||||||
|
"owner",
|
||||||
|
models.ForeignKey(
|
||||||
|
on_delete=django.db.models.deletion.CASCADE,
|
||||||
|
related_name="filter_presets",
|
||||||
|
to=settings.AUTH_USER_MODEL,
|
||||||
|
),
|
||||||
|
),
|
||||||
|
],
|
||||||
|
options={"ordering": ["name", "id"]},
|
||||||
|
),
|
||||||
|
]
|
||||||
@@ -1,6 +1,8 @@
|
|||||||
import decimal
|
import decimal
|
||||||
import logging
|
import logging
|
||||||
|
import uuid
|
||||||
from copy import deepcopy
|
from copy import deepcopy
|
||||||
|
from pathlib import Path
|
||||||
|
|
||||||
from apps.common.fields.month_year import MonthYearModelField
|
from apps.common.fields.month_year import MonthYearModelField
|
||||||
from apps.common.functions.decimals import truncate_decimal
|
from apps.common.functions.decimals import truncate_decimal
|
||||||
@@ -13,26 +15,47 @@ from apps.common.models import (
|
|||||||
)
|
)
|
||||||
from apps.common.templatetags.decimal import drop_trailing_zeros, localize_number
|
from apps.common.templatetags.decimal import drop_trailing_zeros, localize_number
|
||||||
from apps.currencies.utils.convert import convert
|
from apps.currencies.utils.convert import convert
|
||||||
|
from apps.transactions.storage import PrivateMediaStorage
|
||||||
from apps.transactions.validators import validate_decimal_places, validate_non_negative
|
from apps.transactions.validators import validate_decimal_places, validate_non_negative
|
||||||
from dateutil.relativedelta import relativedelta
|
from dateutil.relativedelta import relativedelta
|
||||||
from django.conf import settings
|
from django.conf import settings
|
||||||
from django.core.validators import MinValueValidator
|
from django.core.validators import MinValueValidator
|
||||||
from django.db import models, transaction
|
from django.db import models, transaction
|
||||||
from django.db.models import Q
|
from django.db.models import Q
|
||||||
from django.dispatch import Signal
|
from django.db.models.signals import post_delete
|
||||||
from django.forms import ValidationError
|
from django.dispatch import Signal, receiver
|
||||||
from django.template.defaultfilters import date
|
from django.template.defaultfilters import date
|
||||||
from django.utils import timezone
|
from django.utils import timezone
|
||||||
from django.utils.translation import gettext_lazy as _
|
from django.utils.translation import gettext_lazy as _
|
||||||
|
|
||||||
logger = logging.getLogger()
|
logger = logging.getLogger()
|
||||||
|
|
||||||
|
|
||||||
transaction_created = Signal()
|
transaction_created = Signal()
|
||||||
transaction_updated = Signal()
|
transaction_updated = Signal()
|
||||||
transaction_deleted = Signal()
|
transaction_deleted = Signal()
|
||||||
|
|
||||||
|
|
||||||
|
class FilterPreset(models.Model):
|
||||||
|
owner = models.ForeignKey(
|
||||||
|
settings.AUTH_USER_MODEL,
|
||||||
|
on_delete=models.CASCADE,
|
||||||
|
related_name="filter_presets",
|
||||||
|
)
|
||||||
|
name = models.CharField(max_length=100)
|
||||||
|
parameters = models.JSONField(default=dict)
|
||||||
|
|
||||||
|
class Meta:
|
||||||
|
ordering = ["name", "id"]
|
||||||
|
|
||||||
|
def __str__(self):
|
||||||
|
return self.name
|
||||||
|
|
||||||
|
|
||||||
|
def transaction_attachment_path(instance, filename):
|
||||||
|
extension = Path(filename).suffix
|
||||||
|
return f"transaction_attachments/{instance.transaction_id}/{instance.id}{extension}"
|
||||||
|
|
||||||
|
|
||||||
class SoftDeleteQuerySet(models.QuerySet):
|
class SoftDeleteQuerySet(models.QuerySet):
|
||||||
@staticmethod
|
@staticmethod
|
||||||
def _emit_signals(instances, created=False, old_data=None):
|
def _emit_signals(instances, created=False, old_data=None):
|
||||||
@@ -384,6 +407,10 @@ class Transaction(OwnedObject):
|
|||||||
def clean(self):
|
def clean(self):
|
||||||
super().clean()
|
super().clean()
|
||||||
|
|
||||||
|
# Convert empty internal_id to None to allow multiple "empty" values with unique constraint
|
||||||
|
if self.internal_id == "":
|
||||||
|
self.internal_id = None
|
||||||
|
|
||||||
# Only process amount and reference_date if account exists
|
# Only process amount and reference_date if account exists
|
||||||
# If account is missing, Django's required field validation will handle it
|
# If account is missing, Django's required field validation will handle it
|
||||||
try:
|
try:
|
||||||
@@ -524,6 +551,62 @@ class Transaction(OwnedObject):
|
|||||||
return new_obj
|
return new_obj
|
||||||
|
|
||||||
|
|
||||||
|
class TransactionAttachment(models.Model):
|
||||||
|
id = models.UUIDField(primary_key=True, default=uuid.uuid4, editable=False)
|
||||||
|
transaction = models.ForeignKey(
|
||||||
|
Transaction,
|
||||||
|
on_delete=models.CASCADE,
|
||||||
|
related_name="attachments",
|
||||||
|
verbose_name=_("Transaction"),
|
||||||
|
)
|
||||||
|
file = models.FileField(
|
||||||
|
upload_to=transaction_attachment_path,
|
||||||
|
storage=PrivateMediaStorage(),
|
||||||
|
verbose_name=_("File"),
|
||||||
|
)
|
||||||
|
original_name = models.CharField(max_length=255, verbose_name=_("Original Name"))
|
||||||
|
content_type = models.CharField(
|
||||||
|
max_length=255, blank=True, verbose_name=_("Content Type")
|
||||||
|
)
|
||||||
|
size = models.PositiveBigIntegerField(default=0, verbose_name=_("Size"))
|
||||||
|
uploaded_by = models.ForeignKey(
|
||||||
|
settings.AUTH_USER_MODEL,
|
||||||
|
on_delete=models.PROTECT,
|
||||||
|
related_name="transaction_attachments",
|
||||||
|
verbose_name=_("Uploaded By"),
|
||||||
|
)
|
||||||
|
created_at = models.DateTimeField(auto_now_add=True)
|
||||||
|
|
||||||
|
class Meta:
|
||||||
|
verbose_name = _("Transaction Attachment")
|
||||||
|
verbose_name_plural = _("Transaction Attachments")
|
||||||
|
db_table = "transaction_attachments"
|
||||||
|
ordering = ["-created_at", "original_name"]
|
||||||
|
|
||||||
|
def save(self, *args, **kwargs):
|
||||||
|
if self.file:
|
||||||
|
if not self.original_name:
|
||||||
|
self.original_name = Path(self.file.name).name
|
||||||
|
if not self.size:
|
||||||
|
self.size = self.file.size
|
||||||
|
if not self.content_type:
|
||||||
|
self.content_type = getattr(self.file.file, "content_type", "")
|
||||||
|
super().save(*args, **kwargs)
|
||||||
|
|
||||||
|
def __str__(self):
|
||||||
|
return self.original_name
|
||||||
|
|
||||||
|
|
||||||
|
@receiver(post_delete, sender=TransactionAttachment)
|
||||||
|
def delete_transaction_attachment_file(sender, instance, **kwargs):
|
||||||
|
if not instance.file.name:
|
||||||
|
return
|
||||||
|
|
||||||
|
storage = instance.file.storage
|
||||||
|
if storage.exists(instance.file.name):
|
||||||
|
storage.delete(instance.file.name)
|
||||||
|
|
||||||
|
|
||||||
class InstallmentPlan(models.Model):
|
class InstallmentPlan(models.Model):
|
||||||
class Recurrence(models.TextChoices):
|
class Recurrence(models.TextChoices):
|
||||||
YEARLY = "yearly", _("Yearly")
|
YEARLY = "yearly", _("Yearly")
|
||||||
@@ -871,10 +954,10 @@ class RecurringTransaction(models.Model):
|
|||||||
notes=self.notes if self.add_notes_to_transaction else "",
|
notes=self.notes if self.add_notes_to_transaction else "",
|
||||||
owner=self.account.owner,
|
owner=self.account.owner,
|
||||||
)
|
)
|
||||||
if self.tags.exists():
|
# Unfiltered managers: generation also runs without a current user, or with a
|
||||||
created_transaction.tags.set(self.tags.all())
|
# different one, and the scoped default manager would hide private rows.
|
||||||
if self.entities.exists():
|
created_transaction.tags.set(self.tags(manager="all_objects").all())
|
||||||
created_transaction.entities.set(self.entities.all())
|
created_transaction.entities.set(self.entities(manager="all_objects").all())
|
||||||
|
|
||||||
def get_recurrence_delta(self):
|
def get_recurrence_delta(self):
|
||||||
if self.recurrence_type == self.RecurrenceType.DAY:
|
if self.recurrence_type == self.RecurrenceType.DAY:
|
||||||
@@ -966,9 +1049,11 @@ class RecurringTransaction(models.Model):
|
|||||||
self.notes if self.add_notes_to_transaction else ""
|
self.notes if self.add_notes_to_transaction else ""
|
||||||
)
|
)
|
||||||
|
|
||||||
# Update many-to-many relationships
|
# Update many-to-many relationships (see create_transaction)
|
||||||
existing_transaction.tags.set(self.tags.all())
|
existing_transaction.tags.set(self.tags(manager="all_objects").all())
|
||||||
existing_transaction.entities.set(self.entities.all())
|
existing_transaction.entities.set(
|
||||||
|
self.entities(manager="all_objects").all()
|
||||||
|
)
|
||||||
|
|
||||||
# Save updated transaction
|
# Save updated transaction
|
||||||
existing_transaction.save()
|
existing_transaction.save()
|
||||||
|
|||||||
@@ -0,0 +1,9 @@
|
|||||||
|
from django.conf import settings
|
||||||
|
from django.core.files.storage import FileSystemStorage
|
||||||
|
|
||||||
|
|
||||||
|
class PrivateMediaStorage(FileSystemStorage):
|
||||||
|
def __init__(self, *args, **kwargs):
|
||||||
|
kwargs.setdefault("location", settings.ATTACHMENT_MEDIA_ROOT)
|
||||||
|
kwargs.setdefault("base_url", None)
|
||||||
|
super().__init__(*args, **kwargs)
|
||||||
@@ -0,0 +1,219 @@
|
|||||||
|
import shutil
|
||||||
|
import tempfile
|
||||||
|
from datetime import date
|
||||||
|
from decimal import Decimal
|
||||||
|
from pathlib import Path
|
||||||
|
|
||||||
|
from apps.accounts.models import Account
|
||||||
|
from apps.common.middleware.thread_local import delete_current_user, write_current_user
|
||||||
|
from apps.currencies.models import Currency
|
||||||
|
from apps.transactions.models import Transaction, TransactionAttachment
|
||||||
|
from django.contrib.auth import get_user_model
|
||||||
|
from django.core.files.uploadedfile import SimpleUploadedFile
|
||||||
|
from django.test import TestCase, override_settings
|
||||||
|
from django.urls import reverse
|
||||||
|
|
||||||
|
|
||||||
|
@override_settings(
|
||||||
|
STORAGES={
|
||||||
|
"default": {"BACKEND": "django.core.files.storage.FileSystemStorage"},
|
||||||
|
"staticfiles": {
|
||||||
|
"BACKEND": "django.contrib.staticfiles.storage.StaticFilesStorage"
|
||||||
|
},
|
||||||
|
},
|
||||||
|
WHITENOISE_AUTOREFRESH=True,
|
||||||
|
)
|
||||||
|
class TransactionAttachmentTests(TestCase):
|
||||||
|
def setUp(self):
|
||||||
|
self.attachment_media_root = tempfile.mkdtemp()
|
||||||
|
self.override_private_media = override_settings(
|
||||||
|
ATTACHMENT_MEDIA_ROOT=self.attachment_media_root
|
||||||
|
)
|
||||||
|
self.override_private_media.enable()
|
||||||
|
self.addCleanup(self.override_private_media.disable)
|
||||||
|
self.addCleanup(shutil.rmtree, self.attachment_media_root, ignore_errors=True)
|
||||||
|
|
||||||
|
self.attachment_storage = TransactionAttachment._meta.get_field("file").storage
|
||||||
|
self.original_storage_location = self.attachment_storage._location
|
||||||
|
self.attachment_storage._location = self.attachment_media_root
|
||||||
|
self.attachment_storage.__dict__.pop("base_location", None)
|
||||||
|
self.attachment_storage.__dict__.pop("location", None)
|
||||||
|
self.addCleanup(self.restore_attachment_storage)
|
||||||
|
|
||||||
|
User = get_user_model()
|
||||||
|
self.user1 = User.objects.create_user(
|
||||||
|
email="user1@test.com", password="testpass123"
|
||||||
|
)
|
||||||
|
self.user2 = User.objects.create_user(
|
||||||
|
email="user2@test.com", password="testpass123"
|
||||||
|
)
|
||||||
|
|
||||||
|
self.currency = Currency.objects.create(
|
||||||
|
code="USD", name="US Dollar", decimal_places=2, prefix="$ "
|
||||||
|
)
|
||||||
|
self.user1_account = Account.all_objects.create(
|
||||||
|
name="User1 Account", currency=self.currency, owner=self.user1
|
||||||
|
)
|
||||||
|
self.user2_account = Account.all_objects.create(
|
||||||
|
name="User2 Account", currency=self.currency, owner=self.user2
|
||||||
|
)
|
||||||
|
self.transaction = Transaction.userless_all_objects.create(
|
||||||
|
account=self.user1_account,
|
||||||
|
type=Transaction.Type.EXPENSE,
|
||||||
|
amount=Decimal("12.34"),
|
||||||
|
is_paid=True,
|
||||||
|
date=date(2026, 6, 5),
|
||||||
|
description="Receipt transaction",
|
||||||
|
owner=self.user1,
|
||||||
|
)
|
||||||
|
self.other_transaction = Transaction.userless_all_objects.create(
|
||||||
|
account=self.user2_account,
|
||||||
|
type=Transaction.Type.EXPENSE,
|
||||||
|
amount=Decimal("56.78"),
|
||||||
|
is_paid=True,
|
||||||
|
date=date(2026, 6, 5),
|
||||||
|
description="Other receipt transaction",
|
||||||
|
owner=self.user2,
|
||||||
|
)
|
||||||
|
|
||||||
|
def restore_attachment_storage(self):
|
||||||
|
self.attachment_storage._location = self.original_storage_location
|
||||||
|
self.attachment_storage.__dict__.pop("base_location", None)
|
||||||
|
self.attachment_storage.__dict__.pop("location", None)
|
||||||
|
|
||||||
|
def test_attachment_uses_uuid_and_preserves_original_download_name(self):
|
||||||
|
attachment = TransactionAttachment.objects.create(
|
||||||
|
transaction=self.transaction,
|
||||||
|
file=SimpleUploadedFile(
|
||||||
|
"receipt June.pdf", b"receipt bytes", content_type="application/pdf"
|
||||||
|
),
|
||||||
|
uploaded_by=self.user1,
|
||||||
|
)
|
||||||
|
|
||||||
|
self.assertEqual(attachment.original_name, "receipt June.pdf")
|
||||||
|
self.assertNotIn("receipt June.pdf", attachment.file.name)
|
||||||
|
|
||||||
|
self.client.force_login(self.user1)
|
||||||
|
response = self.client.get(
|
||||||
|
reverse(
|
||||||
|
"transaction_attachment_download",
|
||||||
|
kwargs={"attachment_id": attachment.id},
|
||||||
|
)
|
||||||
|
)
|
||||||
|
|
||||||
|
self.assertEqual(response.status_code, 200)
|
||||||
|
self.assertEqual(b"".join(response.streaming_content), b"receipt bytes")
|
||||||
|
self.assertIn('filename="receipt June.pdf"', response["Content-Disposition"])
|
||||||
|
|
||||||
|
def test_user_without_transaction_access_cannot_download_attachment(self):
|
||||||
|
attachment = TransactionAttachment.objects.create(
|
||||||
|
transaction=self.other_transaction,
|
||||||
|
file=SimpleUploadedFile("private.txt", b"private"),
|
||||||
|
uploaded_by=self.user2,
|
||||||
|
)
|
||||||
|
|
||||||
|
self.client.force_login(self.user1)
|
||||||
|
response = self.client.get(
|
||||||
|
reverse(
|
||||||
|
"transaction_attachment_download",
|
||||||
|
kwargs={"attachment_id": attachment.id},
|
||||||
|
)
|
||||||
|
)
|
||||||
|
|
||||||
|
self.assertEqual(response.status_code, 404)
|
||||||
|
|
||||||
|
def test_attachment_button_lives_in_transaction_hover_toolbar(self):
|
||||||
|
template = Path("templates/cotton/transaction/item.html").read_text()
|
||||||
|
before_toolbar, toolbar = template.split("{# Item actions#}", 1)
|
||||||
|
|
||||||
|
self.assertNotIn("transaction_attachments", before_toolbar)
|
||||||
|
self.assertLess(
|
||||||
|
toolbar.index("transaction_edit"),
|
||||||
|
toolbar.index("transaction_attachments"),
|
||||||
|
)
|
||||||
|
self.assertLess(
|
||||||
|
toolbar.index("transaction_attachments"),
|
||||||
|
toolbar.index("transaction_delete"),
|
||||||
|
)
|
||||||
|
|
||||||
|
def test_transaction_edit_form_does_not_include_attachment_upload(self):
|
||||||
|
self.client.force_login(self.user1)
|
||||||
|
|
||||||
|
response = self.client.get(
|
||||||
|
reverse("transaction_edit", kwargs={"transaction_id": self.transaction.id}),
|
||||||
|
HTTP_HX_REQUEST="true",
|
||||||
|
)
|
||||||
|
|
||||||
|
self.assertEqual(response.status_code, 200)
|
||||||
|
self.assertNotContains(response, "multipart/form-data")
|
||||||
|
self.assertNotContains(response, 'type="file"')
|
||||||
|
|
||||||
|
def test_attachment_management_uploads_multiple_attachments(self):
|
||||||
|
self.client.force_login(self.user1)
|
||||||
|
|
||||||
|
response = self.client.post(
|
||||||
|
reverse(
|
||||||
|
"transaction_attachments",
|
||||||
|
kwargs={"transaction_id": self.transaction.id},
|
||||||
|
),
|
||||||
|
{
|
||||||
|
"attachments": [
|
||||||
|
SimpleUploadedFile("first.txt", b"first"),
|
||||||
|
SimpleUploadedFile("second.txt", b"second"),
|
||||||
|
],
|
||||||
|
},
|
||||||
|
HTTP_HX_REQUEST="true",
|
||||||
|
)
|
||||||
|
|
||||||
|
self.assertEqual(response.status_code, 200)
|
||||||
|
self.assertContains(response, "first.txt")
|
||||||
|
self.assertContains(response, "second.txt")
|
||||||
|
self.assertEqual(self.transaction.attachments.count(), 2)
|
||||||
|
|
||||||
|
def test_attachment_delete_returns_refreshed_attachment_list(self):
|
||||||
|
attachment = TransactionAttachment.objects.create(
|
||||||
|
transaction=self.transaction,
|
||||||
|
file=SimpleUploadedFile("delete-me.txt", b"delete"),
|
||||||
|
uploaded_by=self.user1,
|
||||||
|
)
|
||||||
|
|
||||||
|
self.client.force_login(self.user1)
|
||||||
|
response = self.client.delete(
|
||||||
|
reverse(
|
||||||
|
"transaction_attachment_delete",
|
||||||
|
kwargs={"attachment_id": attachment.id},
|
||||||
|
),
|
||||||
|
HTTP_HX_REQUEST="true",
|
||||||
|
)
|
||||||
|
|
||||||
|
self.assertEqual(response.status_code, 200)
|
||||||
|
self.assertNotContains(response, "delete-me.txt")
|
||||||
|
self.assertContains(response, "No attachments yet")
|
||||||
|
self.assertFalse(
|
||||||
|
TransactionAttachment.objects.filter(id=attachment.id).exists()
|
||||||
|
)
|
||||||
|
|
||||||
|
def test_hard_deleting_transaction_deletes_attachment_files(self):
|
||||||
|
attachment = TransactionAttachment.objects.create(
|
||||||
|
transaction=self.transaction,
|
||||||
|
file=SimpleUploadedFile("hard-delete.txt", b"delete with transaction"),
|
||||||
|
uploaded_by=self.user1,
|
||||||
|
)
|
||||||
|
file_path = Path(attachment.file.path)
|
||||||
|
|
||||||
|
self.assertTrue(file_path.exists())
|
||||||
|
|
||||||
|
write_current_user(self.user1)
|
||||||
|
self.addCleanup(delete_current_user)
|
||||||
|
|
||||||
|
self.transaction.delete()
|
||||||
|
|
||||||
|
self.assertTrue(file_path.exists())
|
||||||
|
self.assertTrue(TransactionAttachment.objects.filter(id=attachment.id).exists())
|
||||||
|
|
||||||
|
self.transaction.delete()
|
||||||
|
|
||||||
|
self.assertFalse(file_path.exists())
|
||||||
|
self.assertFalse(
|
||||||
|
TransactionAttachment.objects.filter(id=attachment.id).exists()
|
||||||
|
)
|
||||||
@@ -0,0 +1,183 @@
|
|||||||
|
from django.contrib.auth import get_user_model
|
||||||
|
from django.test import TestCase, override_settings
|
||||||
|
from django.urls import reverse
|
||||||
|
|
||||||
|
from apps.transactions.models import FilterPreset
|
||||||
|
|
||||||
|
|
||||||
|
@override_settings(
|
||||||
|
STORAGES={
|
||||||
|
"default": {"BACKEND": "django.core.files.storage.FileSystemStorage"},
|
||||||
|
"staticfiles": {
|
||||||
|
"BACKEND": "django.contrib.staticfiles.storage.StaticFilesStorage"
|
||||||
|
},
|
||||||
|
},
|
||||||
|
WHITENOISE_AUTOREFRESH=True,
|
||||||
|
)
|
||||||
|
class FilterPresetViewTests(TestCase):
|
||||||
|
def setUp(self):
|
||||||
|
user_model = get_user_model()
|
||||||
|
self.user = user_model.objects.create_user(
|
||||||
|
email="preset-owner@example.com", password="testpass123"
|
||||||
|
)
|
||||||
|
self.other_user = user_model.objects.create_user(
|
||||||
|
email="other-user@example.com", password="testpass123"
|
||||||
|
)
|
||||||
|
self.client.force_login(self.user)
|
||||||
|
self.preset = FilterPreset.objects.create(
|
||||||
|
owner=self.user,
|
||||||
|
name="Unpaid",
|
||||||
|
parameters={"is_paid": ["0"], "type": ["IN", "EX"]},
|
||||||
|
)
|
||||||
|
|
||||||
|
def test_create_stores_only_transaction_filter_fields(self):
|
||||||
|
response = self.client.post(
|
||||||
|
reverse("filter_preset_create"),
|
||||||
|
{
|
||||||
|
"name": "Account X Unpaid",
|
||||||
|
"account": ["Account X"],
|
||||||
|
"is_paid": ["0"],
|
||||||
|
"order": "newer",
|
||||||
|
},
|
||||||
|
HTTP_HX_REQUEST="true",
|
||||||
|
)
|
||||||
|
|
||||||
|
preset = FilterPreset.objects.get(
|
||||||
|
owner=self.user, name="Account X Unpaid"
|
||||||
|
)
|
||||||
|
self.assertEqual(
|
||||||
|
preset.parameters,
|
||||||
|
{"account": ["Account X"], "is_paid": ["0"]},
|
||||||
|
)
|
||||||
|
self.assertContains(response, "Account X Unpaid")
|
||||||
|
|
||||||
|
def test_create_rejects_a_blank_name(self):
|
||||||
|
response = self.client.post(
|
||||||
|
reverse("filter_preset_create"),
|
||||||
|
{"name": " ", "is_paid": ["0"]},
|
||||||
|
HTTP_HX_REQUEST="true",
|
||||||
|
)
|
||||||
|
|
||||||
|
self.assertEqual(response.status_code, 400)
|
||||||
|
self.assertEqual(FilterPreset.objects.filter(owner=self.user).count(), 1)
|
||||||
|
|
||||||
|
def test_apply_returns_the_saved_filter_form_without_changing_the_url(self):
|
||||||
|
response = self.client.get(
|
||||||
|
reverse("filter_preset_apply", args=[self.preset.pk]),
|
||||||
|
HTTP_HX_REQUEST="true",
|
||||||
|
)
|
||||||
|
|
||||||
|
self.assertEqual(response.status_code, 200)
|
||||||
|
self.assertContains(response, 'id="filter"')
|
||||||
|
self.assertEqual(response.context["filter"].data.getlist("is_paid"), ["0"])
|
||||||
|
self.assertEqual(
|
||||||
|
response.context["filter"].data.getlist("type"), ["IN", "EX"]
|
||||||
|
)
|
||||||
|
self.assertNotIn("HX-Push-Url", response.headers)
|
||||||
|
self.assertNotIn("HX-Replace-Url", response.headers)
|
||||||
|
self.assertIs(response.context.get("filter_is_active"), True)
|
||||||
|
self.assertContains(response, 'hx-swap-oob="outerHTML"')
|
||||||
|
self.assertEqual(
|
||||||
|
response.headers["HX-Trigger-After-Settle"],
|
||||||
|
"updated",
|
||||||
|
)
|
||||||
|
|
||||||
|
def test_clear_returns_the_default_filter_form_without_changing_the_url(self):
|
||||||
|
response = self.client.get(
|
||||||
|
"/transactions/filter/clear/",
|
||||||
|
HTTP_HX_REQUEST="true",
|
||||||
|
)
|
||||||
|
|
||||||
|
self.assertEqual(response.status_code, 200)
|
||||||
|
self.assertContains(response, 'id="filter"')
|
||||||
|
self.assertEqual(
|
||||||
|
response.context["filter"].data.getlist("type"), ["IN", "EX"]
|
||||||
|
)
|
||||||
|
self.assertEqual(
|
||||||
|
response.context["filter"].data.getlist("is_paid"), ["1", "0"]
|
||||||
|
)
|
||||||
|
self.assertNotIn("HX-Push-Url", response.headers)
|
||||||
|
self.assertNotIn("HX-Replace-Url", response.headers)
|
||||||
|
self.assertIs(response.context.get("filter_is_active"), False)
|
||||||
|
self.assertContains(response, 'hx-swap-oob="outerHTML"')
|
||||||
|
self.assertEqual(
|
||||||
|
response.headers["HX-Trigger-After-Settle"],
|
||||||
|
"updated",
|
||||||
|
)
|
||||||
|
|
||||||
|
def test_other_users_cannot_apply_or_delete_a_preset(self):
|
||||||
|
self.client.force_login(self.other_user)
|
||||||
|
|
||||||
|
apply_response = self.client.get(
|
||||||
|
reverse("filter_preset_apply", args=[self.preset.pk]),
|
||||||
|
HTTP_HX_REQUEST="true",
|
||||||
|
)
|
||||||
|
delete_response = self.client.post(
|
||||||
|
reverse("filter_preset_delete", args=[self.preset.pk]),
|
||||||
|
HTTP_HX_REQUEST="true",
|
||||||
|
)
|
||||||
|
|
||||||
|
self.assertEqual(apply_response.status_code, 404)
|
||||||
|
self.assertEqual(delete_response.status_code, 404)
|
||||||
|
self.assertTrue(FilterPreset.objects.filter(pk=self.preset.pk).exists())
|
||||||
|
|
||||||
|
def test_delete_removes_the_current_users_preset(self):
|
||||||
|
response = self.client.post(
|
||||||
|
reverse("filter_preset_delete", args=[self.preset.pk]),
|
||||||
|
HTTP_HX_REQUEST="true",
|
||||||
|
)
|
||||||
|
|
||||||
|
self.assertEqual(response.status_code, 200)
|
||||||
|
self.assertFalse(FilterPreset.objects.filter(pk=self.preset.pk).exists())
|
||||||
|
self.assertNotContains(response, self.preset.name)
|
||||||
|
|
||||||
|
def test_all_transactions_page_only_offers_the_current_users_presets(self):
|
||||||
|
FilterPreset.objects.create(
|
||||||
|
owner=self.other_user,
|
||||||
|
name="Other User Preset",
|
||||||
|
parameters={"type": ["IN"]},
|
||||||
|
)
|
||||||
|
|
||||||
|
response = self.client.get(reverse("transactions_all_index"))
|
||||||
|
|
||||||
|
self.assertContains(response, self.preset.name)
|
||||||
|
self.assertNotContains(response, "Other User Preset")
|
||||||
|
self.assertContains(response, 'id="filter-presets"')
|
||||||
|
self.assertNotContains(response, 'href="./?')
|
||||||
|
self.assertContains(
|
||||||
|
response,
|
||||||
|
reverse("filter_preset_apply", args=[self.preset.pk]),
|
||||||
|
)
|
||||||
|
|
||||||
|
def test_all_transactions_page_marks_a_filtered_query_active(self):
|
||||||
|
response = self.client.get(
|
||||||
|
reverse("transactions_all_index"),
|
||||||
|
{"type": ["IN", "EX"], "is_paid": ["0"]},
|
||||||
|
)
|
||||||
|
|
||||||
|
self.assertIs(response.context["filter_is_active"], True)
|
||||||
|
|
||||||
|
def test_all_transactions_page_does_not_mark_default_query_active(self):
|
||||||
|
response = self.client.get(
|
||||||
|
reverse("transactions_all_index"),
|
||||||
|
{
|
||||||
|
"type": ["IN", "EX"],
|
||||||
|
"is_paid": ["1", "0"],
|
||||||
|
"mute_status": ["active", "muted"],
|
||||||
|
},
|
||||||
|
)
|
||||||
|
|
||||||
|
self.assertIs(response.context["filter_is_active"], False)
|
||||||
|
|
||||||
|
def test_monthly_page_renders_preset_controls(self):
|
||||||
|
other_preset = FilterPreset.objects.create(
|
||||||
|
owner=self.other_user,
|
||||||
|
name="Other Monthly Preset",
|
||||||
|
parameters={"type": ["IN"]},
|
||||||
|
)
|
||||||
|
response = self.client.get(reverse("monthly_overview", args=[8, 2026]))
|
||||||
|
|
||||||
|
self.assertContains(response, 'id="filter-presets"')
|
||||||
|
self.assertNotContains(response, 'href="./?')
|
||||||
|
self.assertContains(response, reverse("filter_preset_create"))
|
||||||
|
self.assertNotContains(response, other_preset.name)
|
||||||
@@ -7,6 +7,7 @@ from django.utils import timezone
|
|||||||
from apps.transactions.models import (
|
from apps.transactions.models import (
|
||||||
TransactionCategory,
|
TransactionCategory,
|
||||||
TransactionTag,
|
TransactionTag,
|
||||||
|
TransactionEntity,
|
||||||
Transaction,
|
Transaction,
|
||||||
InstallmentPlan,
|
InstallmentPlan,
|
||||||
RecurringTransaction,
|
RecurringTransaction,
|
||||||
@@ -125,6 +126,70 @@ class TransactionTests(TestCase):
|
|||||||
datetime.datetime(day=1, month=2, year=2000).date(),
|
datetime.datetime(day=1, month=2, year=2000).date(),
|
||||||
)
|
)
|
||||||
|
|
||||||
|
def test_empty_internal_id_converts_to_none(self):
|
||||||
|
"""Test that empty string internal_id is converted to None"""
|
||||||
|
transaction = Transaction.objects.create(
|
||||||
|
account=self.account,
|
||||||
|
type=Transaction.Type.EXPENSE,
|
||||||
|
date=timezone.now().date(),
|
||||||
|
amount=Decimal("100.00"),
|
||||||
|
description="Test transaction",
|
||||||
|
internal_id="", # Empty string should become None
|
||||||
|
)
|
||||||
|
self.assertIsNone(transaction.internal_id)
|
||||||
|
|
||||||
|
def test_unique_internal_id_works(self):
|
||||||
|
"""Test that unique non-empty internal_id values work correctly"""
|
||||||
|
transaction1 = Transaction.objects.create(
|
||||||
|
account=self.account,
|
||||||
|
type=Transaction.Type.EXPENSE,
|
||||||
|
date=timezone.now().date(),
|
||||||
|
amount=Decimal("100.00"),
|
||||||
|
description="Test transaction 1",
|
||||||
|
internal_id="unique-id-123",
|
||||||
|
)
|
||||||
|
transaction2 = Transaction.objects.create(
|
||||||
|
account=self.account,
|
||||||
|
type=Transaction.Type.EXPENSE,
|
||||||
|
date=timezone.now().date(),
|
||||||
|
amount=Decimal("100.00"),
|
||||||
|
description="Test transaction 2",
|
||||||
|
internal_id="unique-id-456",
|
||||||
|
)
|
||||||
|
self.assertEqual(transaction1.internal_id, "unique-id-123")
|
||||||
|
self.assertEqual(transaction2.internal_id, "unique-id-456")
|
||||||
|
|
||||||
|
def test_multiple_transactions_with_empty_internal_id(self):
|
||||||
|
"""Test that multiple transactions can have empty/None internal_id"""
|
||||||
|
transaction1 = Transaction.objects.create(
|
||||||
|
account=self.account,
|
||||||
|
type=Transaction.Type.EXPENSE,
|
||||||
|
date=timezone.now().date(),
|
||||||
|
amount=Decimal("100.00"),
|
||||||
|
description="Test transaction 1",
|
||||||
|
internal_id="",
|
||||||
|
)
|
||||||
|
transaction2 = Transaction.objects.create(
|
||||||
|
account=self.account,
|
||||||
|
type=Transaction.Type.EXPENSE,
|
||||||
|
date=timezone.now().date(),
|
||||||
|
amount=Decimal("100.00"),
|
||||||
|
description="Test transaction 2",
|
||||||
|
internal_id="",
|
||||||
|
)
|
||||||
|
transaction3 = Transaction.objects.create(
|
||||||
|
account=self.account,
|
||||||
|
type=Transaction.Type.EXPENSE,
|
||||||
|
date=timezone.now().date(),
|
||||||
|
amount=Decimal("100.00"),
|
||||||
|
description="Test transaction 3",
|
||||||
|
internal_id=None,
|
||||||
|
)
|
||||||
|
# All should be saved successfully with None internal_id
|
||||||
|
self.assertIsNone(transaction1.internal_id)
|
||||||
|
self.assertIsNone(transaction2.internal_id)
|
||||||
|
self.assertIsNone(transaction3.internal_id)
|
||||||
|
|
||||||
|
|
||||||
class InstallmentPlanTests(TestCase):
|
class InstallmentPlanTests(TestCase):
|
||||||
def setUp(self):
|
def setUp(self):
|
||||||
@@ -176,3 +241,27 @@ class RecurringTransactionTests(TestCase):
|
|||||||
self.assertFalse(recurring.is_paused)
|
self.assertFalse(recurring.is_paused)
|
||||||
self.assertEqual(recurring.recurrence_interval, 1)
|
self.assertEqual(recurring.recurrence_interval, 1)
|
||||||
self.assertEqual(recurring.account.currency.code, "USD")
|
self.assertEqual(recurring.account.currency.code, "USD")
|
||||||
|
|
||||||
|
def test_generate_upcoming_transactions_keeps_tags_and_entities(self):
|
||||||
|
"""Generation must copy tags/entities even with no current user"""
|
||||||
|
tag = TransactionTag.objects.create(name="Essential")
|
||||||
|
entity = TransactionEntity.objects.create(name="Landlord")
|
||||||
|
recurring = RecurringTransaction.objects.create(
|
||||||
|
account=self.account,
|
||||||
|
type=Transaction.Type.EXPENSE,
|
||||||
|
amount=Decimal("100.00"),
|
||||||
|
description="Monthly Payment",
|
||||||
|
start_date=timezone.now().date(),
|
||||||
|
recurrence_type=RecurringTransaction.RecurrenceType.MONTH,
|
||||||
|
recurrence_interval=1,
|
||||||
|
)
|
||||||
|
recurring.tags.set([tag])
|
||||||
|
recurring.entities.set([entity])
|
||||||
|
|
||||||
|
RecurringTransaction.generate_upcoming_transactions()
|
||||||
|
|
||||||
|
generated = Transaction.all_objects.filter(recurring_transaction=recurring)
|
||||||
|
self.assertTrue(generated.exists())
|
||||||
|
for transaction in generated:
|
||||||
|
self.assertIn(tag, transaction.tags(manager="all_objects").all())
|
||||||
|
self.assertIn(entity, transaction.entities(manager="all_objects").all())
|
||||||
|
|||||||
@@ -6,6 +6,26 @@ urlpatterns = [
|
|||||||
path(
|
path(
|
||||||
"transactions/list/", views.transaction_all_list, name="transactions_all_list"
|
"transactions/list/", views.transaction_all_list, name="transactions_all_list"
|
||||||
),
|
),
|
||||||
|
path(
|
||||||
|
"transactions/filter-presets/create/",
|
||||||
|
views.filter_preset_create,
|
||||||
|
name="filter_preset_create",
|
||||||
|
),
|
||||||
|
path(
|
||||||
|
"transactions/filter-presets/<int:preset_id>/apply/",
|
||||||
|
views.filter_preset_apply,
|
||||||
|
name="filter_preset_apply",
|
||||||
|
),
|
||||||
|
path(
|
||||||
|
"transactions/filter-presets/<int:preset_id>/delete/",
|
||||||
|
views.filter_preset_delete,
|
||||||
|
name="filter_preset_delete",
|
||||||
|
),
|
||||||
|
path(
|
||||||
|
"transactions/filter/clear/",
|
||||||
|
views.transaction_filter_clear,
|
||||||
|
name="transaction_filter_clear",
|
||||||
|
),
|
||||||
path(
|
path(
|
||||||
"transactions/trash/",
|
"transactions/trash/",
|
||||||
views.transactions_trash_can_index,
|
views.transactions_trash_can_index,
|
||||||
@@ -81,6 +101,26 @@ urlpatterns = [
|
|||||||
views.transaction_move_to_today,
|
views.transaction_move_to_today,
|
||||||
name="transaction_move_to_today",
|
name="transaction_move_to_today",
|
||||||
),
|
),
|
||||||
|
path(
|
||||||
|
"transaction/<int:transaction_id>/attachments/",
|
||||||
|
views.transaction_attachments,
|
||||||
|
name="transaction_attachments",
|
||||||
|
),
|
||||||
|
path(
|
||||||
|
"transaction/<int:transaction_id>/attachments/list/",
|
||||||
|
views.transaction_attachments_list,
|
||||||
|
name="transaction_attachments_list",
|
||||||
|
),
|
||||||
|
path(
|
||||||
|
"transaction/attachments/<uuid:attachment_id>/download/",
|
||||||
|
views.transaction_attachment_download,
|
||||||
|
name="transaction_attachment_download",
|
||||||
|
),
|
||||||
|
path(
|
||||||
|
"transaction/attachments/<uuid:attachment_id>/delete/",
|
||||||
|
views.transaction_attachment_delete,
|
||||||
|
name="transaction_attachment_delete",
|
||||||
|
),
|
||||||
path(
|
path(
|
||||||
"transaction/<int:transaction_id>/delete/",
|
"transaction/<int:transaction_id>/delete/",
|
||||||
views.transaction_delete,
|
views.transaction_delete,
|
||||||
|
|||||||
@@ -1,11 +1,17 @@
|
|||||||
from django.contrib import messages
|
from django.contrib import messages
|
||||||
from django.contrib.auth.decorators import login_required
|
from django.contrib.auth.decorators import login_required
|
||||||
|
from django.core.exceptions import PermissionDenied
|
||||||
from django.http import HttpResponse
|
from django.http import HttpResponse
|
||||||
from django.shortcuts import render, get_object_or_404
|
from django.shortcuts import render
|
||||||
from django.utils.translation import gettext_lazy as _
|
from django.utils.translation import gettext_lazy as _
|
||||||
from django.views.decorators.http import require_http_methods
|
from django.views.decorators.http import require_http_methods
|
||||||
|
|
||||||
from apps.common.decorators.htmx import only_htmx
|
from apps.common.decorators.htmx import only_htmx
|
||||||
|
from apps.common.functions.permissions import (
|
||||||
|
EDIT,
|
||||||
|
READ,
|
||||||
|
get_shared_object_or_error,
|
||||||
|
)
|
||||||
from apps.transactions.forms import TransactionCategoryForm
|
from apps.transactions.forms import TransactionCategoryForm
|
||||||
from apps.transactions.models import TransactionCategory
|
from apps.transactions.models import TransactionCategory
|
||||||
from apps.common.models import SharedObject
|
from apps.common.models import SharedObject
|
||||||
@@ -35,7 +41,7 @@ def categories_list(request):
|
|||||||
@login_required
|
@login_required
|
||||||
@require_http_methods(["GET"])
|
@require_http_methods(["GET"])
|
||||||
def categories_table_active(request):
|
def categories_table_active(request):
|
||||||
categories = TransactionCategory.objects.filter(active=True).order_by("id")
|
categories = TransactionCategory.objects.filter(active=True).order_by("name")
|
||||||
return render(
|
return render(
|
||||||
request,
|
request,
|
||||||
"categories/fragments/table.html",
|
"categories/fragments/table.html",
|
||||||
@@ -47,7 +53,7 @@ def categories_table_active(request):
|
|||||||
@login_required
|
@login_required
|
||||||
@require_http_methods(["GET"])
|
@require_http_methods(["GET"])
|
||||||
def categories_table_archived(request):
|
def categories_table_archived(request):
|
||||||
categories = TransactionCategory.objects.filter(active=False).order_by("id")
|
categories = TransactionCategory.objects.filter(active=False).order_by("name")
|
||||||
return render(
|
return render(
|
||||||
request,
|
request,
|
||||||
"categories/fragments/table.html",
|
"categories/fragments/table.html",
|
||||||
@@ -85,17 +91,9 @@ def category_add(request, **kwargs):
|
|||||||
@login_required
|
@login_required
|
||||||
@require_http_methods(["GET", "POST"])
|
@require_http_methods(["GET", "POST"])
|
||||||
def category_edit(request, category_id):
|
def category_edit(request, category_id):
|
||||||
category = get_object_or_404(TransactionCategory, id=category_id)
|
category = get_shared_object_or_error(
|
||||||
|
TransactionCategory, request, id=category_id, level=EDIT
|
||||||
if category.owner and category.owner != request.user:
|
)
|
||||||
messages.error(request, _("Only the owner can edit this"))
|
|
||||||
|
|
||||||
return HttpResponse(
|
|
||||||
status=204,
|
|
||||||
headers={
|
|
||||||
"HX-Trigger": "updated, hide_offcanvas",
|
|
||||||
},
|
|
||||||
)
|
|
||||||
|
|
||||||
if request.method == "POST":
|
if request.method == "POST":
|
||||||
form = TransactionCategoryForm(request.POST, instance=category)
|
form = TransactionCategoryForm(request.POST, instance=category)
|
||||||
@@ -123,17 +121,7 @@ def category_edit(request, category_id):
|
|||||||
@login_required
|
@login_required
|
||||||
@require_http_methods(["GET", "POST"])
|
@require_http_methods(["GET", "POST"])
|
||||||
def category_share(request, pk):
|
def category_share(request, pk):
|
||||||
obj = get_object_or_404(TransactionCategory, id=pk)
|
obj = get_shared_object_or_error(TransactionCategory, request, id=pk, level=EDIT)
|
||||||
|
|
||||||
if obj.owner and obj.owner != request.user:
|
|
||||||
messages.error(request, _("Only the owner can edit this"))
|
|
||||||
|
|
||||||
return HttpResponse(
|
|
||||||
status=204,
|
|
||||||
headers={
|
|
||||||
"HX-Trigger": "updated, hide_offcanvas",
|
|
||||||
},
|
|
||||||
)
|
|
||||||
|
|
||||||
if request.method == "POST":
|
if request.method == "POST":
|
||||||
form = SharedObjectForm(request.POST, instance=obj, user=request.user)
|
form = SharedObjectForm(request.POST, instance=obj, user=request.user)
|
||||||
@@ -161,14 +149,20 @@ def category_share(request, pk):
|
|||||||
@login_required
|
@login_required
|
||||||
@require_http_methods(["DELETE"])
|
@require_http_methods(["DELETE"])
|
||||||
def category_delete(request, category_id):
|
def category_delete(request, category_id):
|
||||||
category = get_object_or_404(TransactionCategory, id=category_id)
|
category = get_shared_object_or_error(
|
||||||
|
TransactionCategory, request, id=category_id, level=READ
|
||||||
|
)
|
||||||
|
|
||||||
if category.owner != request.user and request.user in category.shared_with.all():
|
if category.is_editable_by(request.user):
|
||||||
|
category.delete()
|
||||||
|
messages.success(request, _("Category deleted successfully"))
|
||||||
|
elif category.shared_with.filter(pk=request.user.pk).exists():
|
||||||
|
# Someone else's object shared with us: we can drop our own access
|
||||||
|
# to it, but never delete it.
|
||||||
category.shared_with.remove(request.user)
|
category.shared_with.remove(request.user)
|
||||||
messages.success(request, _("Item no longer shared with you"))
|
messages.success(request, _("Item no longer shared with you"))
|
||||||
else:
|
else:
|
||||||
category.delete()
|
raise PermissionDenied
|
||||||
messages.success(request, _("Category deleted successfully"))
|
|
||||||
|
|
||||||
return HttpResponse(
|
return HttpResponse(
|
||||||
status=204,
|
status=204,
|
||||||
@@ -182,7 +176,9 @@ def category_delete(request, category_id):
|
|||||||
@login_required
|
@login_required
|
||||||
@require_http_methods(["GET"])
|
@require_http_methods(["GET"])
|
||||||
def category_take_ownership(request, category_id):
|
def category_take_ownership(request, category_id):
|
||||||
category = get_object_or_404(TransactionCategory, id=category_id)
|
category = get_shared_object_or_error(
|
||||||
|
TransactionCategory, request, id=category_id, level=EDIT
|
||||||
|
)
|
||||||
|
|
||||||
if not category.owner:
|
if not category.owner:
|
||||||
category.owner = request.user
|
category.owner = request.user
|
||||||
|
|||||||
@@ -1,11 +1,17 @@
|
|||||||
from django.contrib import messages
|
from django.contrib import messages
|
||||||
from django.contrib.auth.decorators import login_required
|
from django.contrib.auth.decorators import login_required
|
||||||
|
from django.core.exceptions import PermissionDenied
|
||||||
from django.http import HttpResponse
|
from django.http import HttpResponse
|
||||||
from django.shortcuts import render, get_object_or_404
|
from django.shortcuts import render
|
||||||
from django.utils.translation import gettext_lazy as _
|
from django.utils.translation import gettext_lazy as _
|
||||||
from django.views.decorators.http import require_http_methods
|
from django.views.decorators.http import require_http_methods
|
||||||
|
|
||||||
from apps.common.decorators.htmx import only_htmx
|
from apps.common.decorators.htmx import only_htmx
|
||||||
|
from apps.common.functions.permissions import (
|
||||||
|
EDIT,
|
||||||
|
READ,
|
||||||
|
get_shared_object_or_error,
|
||||||
|
)
|
||||||
from apps.transactions.forms import TransactionEntityForm
|
from apps.transactions.forms import TransactionEntityForm
|
||||||
from apps.transactions.models import TransactionEntity
|
from apps.transactions.models import TransactionEntity
|
||||||
from apps.common.models import SharedObject
|
from apps.common.models import SharedObject
|
||||||
@@ -35,7 +41,7 @@ def entities_list(request):
|
|||||||
@login_required
|
@login_required
|
||||||
@require_http_methods(["GET"])
|
@require_http_methods(["GET"])
|
||||||
def entities_table_active(request):
|
def entities_table_active(request):
|
||||||
entities = TransactionEntity.objects.filter(active=True).order_by("id")
|
entities = TransactionEntity.objects.filter(active=True).order_by("name")
|
||||||
return render(
|
return render(
|
||||||
request,
|
request,
|
||||||
"entities/fragments/table.html",
|
"entities/fragments/table.html",
|
||||||
@@ -47,7 +53,7 @@ def entities_table_active(request):
|
|||||||
@login_required
|
@login_required
|
||||||
@require_http_methods(["GET"])
|
@require_http_methods(["GET"])
|
||||||
def entities_table_archived(request):
|
def entities_table_archived(request):
|
||||||
entities = TransactionEntity.objects.filter(active=False).order_by("id")
|
entities = TransactionEntity.objects.filter(active=False).order_by("name")
|
||||||
return render(
|
return render(
|
||||||
request,
|
request,
|
||||||
"entities/fragments/table.html",
|
"entities/fragments/table.html",
|
||||||
@@ -85,17 +91,9 @@ def entity_add(request, **kwargs):
|
|||||||
@login_required
|
@login_required
|
||||||
@require_http_methods(["GET", "POST"])
|
@require_http_methods(["GET", "POST"])
|
||||||
def entity_edit(request, entity_id):
|
def entity_edit(request, entity_id):
|
||||||
entity = get_object_or_404(TransactionEntity, id=entity_id)
|
entity = get_shared_object_or_error(
|
||||||
|
TransactionEntity, request, id=entity_id, level=EDIT
|
||||||
if entity.owner and entity.owner != request.user:
|
)
|
||||||
messages.error(request, _("Only the owner can edit this"))
|
|
||||||
|
|
||||||
return HttpResponse(
|
|
||||||
status=204,
|
|
||||||
headers={
|
|
||||||
"HX-Trigger": "updated, hide_offcanvas",
|
|
||||||
},
|
|
||||||
)
|
|
||||||
|
|
||||||
if request.method == "POST":
|
if request.method == "POST":
|
||||||
form = TransactionEntityForm(request.POST, instance=entity)
|
form = TransactionEntityForm(request.POST, instance=entity)
|
||||||
@@ -123,14 +121,20 @@ def entity_edit(request, entity_id):
|
|||||||
@login_required
|
@login_required
|
||||||
@require_http_methods(["DELETE"])
|
@require_http_methods(["DELETE"])
|
||||||
def entity_delete(request, entity_id):
|
def entity_delete(request, entity_id):
|
||||||
entity = get_object_or_404(TransactionEntity, id=entity_id)
|
entity = get_shared_object_or_error(
|
||||||
|
TransactionEntity, request, id=entity_id, level=READ
|
||||||
|
)
|
||||||
|
|
||||||
if entity.owner != request.user and request.user in entity.shared_with.all():
|
if entity.is_editable_by(request.user):
|
||||||
|
entity.delete()
|
||||||
|
messages.success(request, _("Entity deleted successfully"))
|
||||||
|
elif entity.shared_with.filter(pk=request.user.pk).exists():
|
||||||
|
# Someone else's object shared with us: we can drop our own access
|
||||||
|
# to it, but never delete it.
|
||||||
entity.shared_with.remove(request.user)
|
entity.shared_with.remove(request.user)
|
||||||
messages.success(request, _("Item no longer shared with you"))
|
messages.success(request, _("Item no longer shared with you"))
|
||||||
else:
|
else:
|
||||||
entity.delete()
|
raise PermissionDenied
|
||||||
messages.success(request, _("Entity deleted successfully"))
|
|
||||||
|
|
||||||
return HttpResponse(
|
return HttpResponse(
|
||||||
status=204,
|
status=204,
|
||||||
@@ -144,7 +148,9 @@ def entity_delete(request, entity_id):
|
|||||||
@login_required
|
@login_required
|
||||||
@require_http_methods(["GET"])
|
@require_http_methods(["GET"])
|
||||||
def entity_take_ownership(request, entity_id):
|
def entity_take_ownership(request, entity_id):
|
||||||
entity = get_object_or_404(TransactionEntity, id=entity_id)
|
entity = get_shared_object_or_error(
|
||||||
|
TransactionEntity, request, id=entity_id, level=EDIT
|
||||||
|
)
|
||||||
|
|
||||||
if not entity.owner:
|
if not entity.owner:
|
||||||
entity.owner = request.user
|
entity.owner = request.user
|
||||||
@@ -165,17 +171,7 @@ def entity_take_ownership(request, entity_id):
|
|||||||
@login_required
|
@login_required
|
||||||
@require_http_methods(["GET", "POST"])
|
@require_http_methods(["GET", "POST"])
|
||||||
def entity_share(request, pk):
|
def entity_share(request, pk):
|
||||||
obj = get_object_or_404(TransactionEntity, id=pk)
|
obj = get_shared_object_or_error(TransactionEntity, request, id=pk, level=EDIT)
|
||||||
|
|
||||||
if obj.owner and obj.owner != request.user:
|
|
||||||
messages.error(request, _("Only the owner can edit this"))
|
|
||||||
|
|
||||||
return HttpResponse(
|
|
||||||
status=204,
|
|
||||||
headers={
|
|
||||||
"HX-Trigger": "updated, hide_offcanvas",
|
|
||||||
},
|
|
||||||
)
|
|
||||||
|
|
||||||
if request.method == "POST":
|
if request.method == "POST":
|
||||||
form = SharedObjectForm(request.POST, instance=obj, user=request.user)
|
form = SharedObjectForm(request.POST, instance=obj, user=request.user)
|
||||||
|
|||||||
@@ -1,11 +1,17 @@
|
|||||||
from django.contrib import messages
|
from django.contrib import messages
|
||||||
from django.contrib.auth.decorators import login_required
|
from django.contrib.auth.decorators import login_required
|
||||||
|
from django.core.exceptions import PermissionDenied
|
||||||
from django.http import HttpResponse
|
from django.http import HttpResponse
|
||||||
from django.shortcuts import render, get_object_or_404
|
from django.shortcuts import render
|
||||||
from django.utils.translation import gettext_lazy as _
|
from django.utils.translation import gettext_lazy as _
|
||||||
from django.views.decorators.http import require_http_methods
|
from django.views.decorators.http import require_http_methods
|
||||||
|
|
||||||
from apps.common.decorators.htmx import only_htmx
|
from apps.common.decorators.htmx import only_htmx
|
||||||
|
from apps.common.functions.permissions import (
|
||||||
|
EDIT,
|
||||||
|
READ,
|
||||||
|
get_shared_object_or_error,
|
||||||
|
)
|
||||||
from apps.transactions.forms import TransactionTagForm
|
from apps.transactions.forms import TransactionTagForm
|
||||||
from apps.transactions.models import TransactionTag
|
from apps.transactions.models import TransactionTag
|
||||||
from apps.common.models import SharedObject
|
from apps.common.models import SharedObject
|
||||||
@@ -35,7 +41,7 @@ def tags_list(request):
|
|||||||
@login_required
|
@login_required
|
||||||
@require_http_methods(["GET"])
|
@require_http_methods(["GET"])
|
||||||
def tags_table_active(request):
|
def tags_table_active(request):
|
||||||
tags = TransactionTag.objects.filter(active=True).order_by("id")
|
tags = TransactionTag.objects.filter(active=True).order_by("name")
|
||||||
return render(
|
return render(
|
||||||
request,
|
request,
|
||||||
"tags/fragments/table.html",
|
"tags/fragments/table.html",
|
||||||
@@ -47,7 +53,7 @@ def tags_table_active(request):
|
|||||||
@login_required
|
@login_required
|
||||||
@require_http_methods(["GET"])
|
@require_http_methods(["GET"])
|
||||||
def tags_table_archived(request):
|
def tags_table_archived(request):
|
||||||
tags = TransactionTag.objects.filter(active=False).order_by("id")
|
tags = TransactionTag.objects.filter(active=False).order_by("name")
|
||||||
return render(
|
return render(
|
||||||
request,
|
request,
|
||||||
"tags/fragments/table.html",
|
"tags/fragments/table.html",
|
||||||
@@ -85,17 +91,7 @@ def tag_add(request, **kwargs):
|
|||||||
@login_required
|
@login_required
|
||||||
@require_http_methods(["GET", "POST"])
|
@require_http_methods(["GET", "POST"])
|
||||||
def tag_edit(request, tag_id):
|
def tag_edit(request, tag_id):
|
||||||
tag = get_object_or_404(TransactionTag, id=tag_id)
|
tag = get_shared_object_or_error(TransactionTag, request, id=tag_id, level=EDIT)
|
||||||
|
|
||||||
if tag.owner and tag.owner != request.user:
|
|
||||||
messages.error(request, _("Only the owner can edit this"))
|
|
||||||
|
|
||||||
return HttpResponse(
|
|
||||||
status=204,
|
|
||||||
headers={
|
|
||||||
"HX-Trigger": "updated, hide_offcanvas",
|
|
||||||
},
|
|
||||||
)
|
|
||||||
|
|
||||||
if request.method == "POST":
|
if request.method == "POST":
|
||||||
form = TransactionTagForm(request.POST, instance=tag)
|
form = TransactionTagForm(request.POST, instance=tag)
|
||||||
@@ -123,14 +119,18 @@ def tag_edit(request, tag_id):
|
|||||||
@login_required
|
@login_required
|
||||||
@require_http_methods(["DELETE"])
|
@require_http_methods(["DELETE"])
|
||||||
def tag_delete(request, tag_id):
|
def tag_delete(request, tag_id):
|
||||||
tag = get_object_or_404(TransactionTag, id=tag_id)
|
tag = get_shared_object_or_error(TransactionTag, request, id=tag_id, level=READ)
|
||||||
|
|
||||||
if tag.owner != request.user and request.user in tag.shared_with.all():
|
if tag.is_editable_by(request.user):
|
||||||
|
tag.delete()
|
||||||
|
messages.success(request, _("Tag deleted successfully"))
|
||||||
|
elif tag.shared_with.filter(pk=request.user.pk).exists():
|
||||||
|
# Someone else's object shared with us: we can drop our own access
|
||||||
|
# to it, but never delete it.
|
||||||
tag.shared_with.remove(request.user)
|
tag.shared_with.remove(request.user)
|
||||||
messages.success(request, _("Item no longer shared with you"))
|
messages.success(request, _("Item no longer shared with you"))
|
||||||
else:
|
else:
|
||||||
tag.delete()
|
raise PermissionDenied
|
||||||
messages.success(request, _("Tag deleted successfully"))
|
|
||||||
|
|
||||||
return HttpResponse(
|
return HttpResponse(
|
||||||
status=204,
|
status=204,
|
||||||
@@ -144,7 +144,7 @@ def tag_delete(request, tag_id):
|
|||||||
@login_required
|
@login_required
|
||||||
@require_http_methods(["GET"])
|
@require_http_methods(["GET"])
|
||||||
def tag_take_ownership(request, tag_id):
|
def tag_take_ownership(request, tag_id):
|
||||||
tag = get_object_or_404(TransactionTag, id=tag_id)
|
tag = get_shared_object_or_error(TransactionTag, request, id=tag_id, level=EDIT)
|
||||||
|
|
||||||
if not tag.owner:
|
if not tag.owner:
|
||||||
tag.owner = request.user
|
tag.owner = request.user
|
||||||
@@ -165,17 +165,7 @@ def tag_take_ownership(request, tag_id):
|
|||||||
@login_required
|
@login_required
|
||||||
@require_http_methods(["GET", "POST"])
|
@require_http_methods(["GET", "POST"])
|
||||||
def tag_share(request, pk):
|
def tag_share(request, pk):
|
||||||
obj = get_object_or_404(TransactionTag, id=pk)
|
obj = get_shared_object_or_error(TransactionTag, request, id=pk, level=EDIT)
|
||||||
|
|
||||||
if obj.owner and obj.owner != request.user:
|
|
||||||
messages.error(request, _("Only the owner can edit this"))
|
|
||||||
|
|
||||||
return HttpResponse(
|
|
||||||
status=204,
|
|
||||||
headers={
|
|
||||||
"HX-Trigger": "updated, hide_offcanvas",
|
|
||||||
},
|
|
||||||
)
|
|
||||||
|
|
||||||
if request.method == "POST":
|
if request.method == "POST":
|
||||||
form = SharedObjectForm(request.POST, instance=obj, user=request.user)
|
form = SharedObjectForm(request.POST, instance=obj, user=request.user)
|
||||||
|
|||||||
@@ -1,32 +1,120 @@
|
|||||||
import datetime
|
import datetime
|
||||||
from copy import deepcopy
|
from copy import deepcopy
|
||||||
|
|
||||||
from dateutil.relativedelta import relativedelta
|
from apps.common.decorators.demo import disabled_on_demo
|
||||||
from django.contrib import messages
|
|
||||||
from django.contrib.auth.decorators import login_required
|
|
||||||
from django.core.paginator import Paginator
|
|
||||||
from django.db.models import Q, When, Case, Value, IntegerField
|
|
||||||
from django.http import HttpResponse, JsonResponse
|
|
||||||
from django.shortcuts import render, get_object_or_404
|
|
||||||
from django.utils import timezone
|
|
||||||
from django.utils.translation import gettext_lazy as _, ngettext_lazy
|
|
||||||
from django.views.decorators.http import require_http_methods
|
|
||||||
|
|
||||||
from apps.common.decorators.htmx import only_htmx
|
from apps.common.decorators.htmx import only_htmx
|
||||||
from apps.rules.signals import transaction_created, transaction_updated
|
from apps.rules.signals import transaction_created, transaction_updated
|
||||||
from apps.transactions.filters import TransactionsFilter
|
from apps.transactions.filters import TransactionsFilter
|
||||||
from apps.transactions.forms import (
|
from apps.transactions.forms import (
|
||||||
|
BulkEditTransactionForm,
|
||||||
|
TransactionAttachmentForm,
|
||||||
TransactionForm,
|
TransactionForm,
|
||||||
TransferForm,
|
TransferForm,
|
||||||
BulkEditTransactionForm,
|
|
||||||
)
|
)
|
||||||
from apps.transactions.models import Transaction
|
from apps.transactions.models import FilterPreset, Transaction, TransactionAttachment
|
||||||
from apps.transactions.utils.calculations import (
|
from apps.transactions.utils.calculations import (
|
||||||
calculate_currency_totals,
|
|
||||||
calculate_account_totals,
|
calculate_account_totals,
|
||||||
|
calculate_currency_totals,
|
||||||
calculate_percentage_distribution,
|
calculate_percentage_distribution,
|
||||||
)
|
)
|
||||||
from apps.transactions.utils.default_ordering import default_order
|
from apps.transactions.utils.default_ordering import default_order
|
||||||
|
from dateutil.relativedelta import relativedelta
|
||||||
|
from django.contrib import messages
|
||||||
|
from django.contrib.auth.decorators import login_required
|
||||||
|
from django.core.paginator import Paginator
|
||||||
|
from django.db.models import Case, IntegerField, Q, Value, When
|
||||||
|
from django.http import FileResponse, Http404, HttpResponse, JsonResponse, QueryDict
|
||||||
|
from django.shortcuts import get_object_or_404, render
|
||||||
|
from django.utils import timezone
|
||||||
|
from django.utils.translation import gettext_lazy as _
|
||||||
|
from django.utils.translation import ngettext_lazy
|
||||||
|
from django.views.decorators.http import require_http_methods
|
||||||
|
|
||||||
|
|
||||||
|
def _get_accessible_transaction_or_404(transaction_id):
|
||||||
|
return get_object_or_404(Transaction.objects, id=transaction_id)
|
||||||
|
|
||||||
|
|
||||||
|
def _get_accessible_attachment_or_404(attachment_id):
|
||||||
|
attachment = get_object_or_404(
|
||||||
|
TransactionAttachment.objects.select_related("transaction"),
|
||||||
|
id=attachment_id,
|
||||||
|
)
|
||||||
|
if not Transaction.objects.filter(id=attachment.transaction_id).exists():
|
||||||
|
raise Http404()
|
||||||
|
return attachment
|
||||||
|
|
||||||
|
|
||||||
|
@only_htmx
|
||||||
|
@login_required
|
||||||
|
@disabled_on_demo
|
||||||
|
@require_http_methods(["GET", "POST"])
|
||||||
|
def transaction_attachments(request, transaction_id):
|
||||||
|
transaction = _get_accessible_transaction_or_404(transaction_id)
|
||||||
|
|
||||||
|
if request.method == "POST":
|
||||||
|
form = TransactionAttachmentForm(request.POST, request.FILES)
|
||||||
|
if form.is_valid():
|
||||||
|
form.save(transaction=transaction, uploaded_by=request.user)
|
||||||
|
messages.success(request, _("Attachment uploaded successfully"))
|
||||||
|
form = TransactionAttachmentForm()
|
||||||
|
else:
|
||||||
|
form = TransactionAttachmentForm()
|
||||||
|
|
||||||
|
response = render(
|
||||||
|
request,
|
||||||
|
"transactions/fragments/attachments_manage.html",
|
||||||
|
{"form": form, "transaction": transaction},
|
||||||
|
)
|
||||||
|
|
||||||
|
response["HX-Trigger"] = "toasts, updated"
|
||||||
|
|
||||||
|
return response
|
||||||
|
|
||||||
|
|
||||||
|
@only_htmx
|
||||||
|
@login_required
|
||||||
|
@disabled_on_demo
|
||||||
|
@require_http_methods(["GET"])
|
||||||
|
def transaction_attachments_list(request, transaction_id):
|
||||||
|
transaction = _get_accessible_transaction_or_404(transaction_id)
|
||||||
|
return render(
|
||||||
|
request,
|
||||||
|
"transactions/fragments/attachments.html",
|
||||||
|
{"transaction": transaction},
|
||||||
|
)
|
||||||
|
|
||||||
|
|
||||||
|
@login_required
|
||||||
|
@disabled_on_demo
|
||||||
|
@require_http_methods(["GET"])
|
||||||
|
def transaction_attachment_download(request, attachment_id):
|
||||||
|
attachment = _get_accessible_attachment_or_404(attachment_id)
|
||||||
|
return FileResponse(
|
||||||
|
attachment.file.open("rb"),
|
||||||
|
as_attachment=False,
|
||||||
|
filename=attachment.original_name,
|
||||||
|
content_type=attachment.content_type or "application/octet-stream",
|
||||||
|
)
|
||||||
|
|
||||||
|
|
||||||
|
@only_htmx
|
||||||
|
@login_required
|
||||||
|
@disabled_on_demo
|
||||||
|
@require_http_methods(["DELETE"])
|
||||||
|
def transaction_attachment_delete(request, attachment_id):
|
||||||
|
attachment = _get_accessible_attachment_or_404(attachment_id)
|
||||||
|
transaction = attachment.transaction
|
||||||
|
attachment.file.delete(save=False)
|
||||||
|
attachment.delete()
|
||||||
|
messages.success(request, _("Attachment deleted successfully"))
|
||||||
|
response = render(
|
||||||
|
request,
|
||||||
|
"transactions/fragments/attachments.html",
|
||||||
|
{"transaction": transaction},
|
||||||
|
)
|
||||||
|
response["HX-Trigger"] = "toasts, updated"
|
||||||
|
return response
|
||||||
|
|
||||||
|
|
||||||
@only_htmx
|
@only_htmx
|
||||||
@@ -152,7 +240,9 @@ def transaction_simple_add(request):
|
|||||||
date_param = request.GET.get("date")
|
date_param = request.GET.get("date")
|
||||||
if date_param:
|
if date_param:
|
||||||
try:
|
try:
|
||||||
initial_data["date"] = datetime.datetime.strptime(date_param, "%Y-%m-%d").date()
|
initial_data["date"] = datetime.datetime.strptime(
|
||||||
|
date_param, "%Y-%m-%d"
|
||||||
|
).date()
|
||||||
except ValueError:
|
except ValueError:
|
||||||
pass
|
pass
|
||||||
|
|
||||||
@@ -160,7 +250,9 @@ def transaction_simple_add(request):
|
|||||||
reference_date_param = request.GET.get("reference_date")
|
reference_date_param = request.GET.get("reference_date")
|
||||||
if reference_date_param:
|
if reference_date_param:
|
||||||
try:
|
try:
|
||||||
initial_data["reference_date"] = datetime.datetime.strptime(reference_date_param, "%Y-%m-%d").date()
|
initial_data["reference_date"] = datetime.datetime.strptime(
|
||||||
|
reference_date_param, "%Y-%m-%d"
|
||||||
|
).date()
|
||||||
except ValueError:
|
except ValueError:
|
||||||
pass
|
pass
|
||||||
|
|
||||||
@@ -172,7 +264,10 @@ def transaction_simple_add(request):
|
|||||||
except (ValueError, TypeError):
|
except (ValueError, TypeError):
|
||||||
# Try to find by name
|
# Try to find by name
|
||||||
from apps.accounts.models import Account
|
from apps.accounts.models import Account
|
||||||
account = Account.objects.filter(name__iexact=account_param, is_archived=False).first()
|
|
||||||
|
account = Account.objects.filter(
|
||||||
|
name__iexact=account_param, is_archived=False
|
||||||
|
).first()
|
||||||
if account:
|
if account:
|
||||||
initial_data["account"] = account.pk
|
initial_data["account"] = account.pk
|
||||||
|
|
||||||
@@ -207,7 +302,10 @@ def transaction_simple_add(request):
|
|||||||
except (ValueError, TypeError):
|
except (ValueError, TypeError):
|
||||||
# Try to find by name
|
# Try to find by name
|
||||||
from apps.transactions.models import TransactionCategory
|
from apps.transactions.models import TransactionCategory
|
||||||
category = TransactionCategory.objects.filter(name__iexact=category_param, active=True).first()
|
|
||||||
|
category = TransactionCategory.objects.filter(
|
||||||
|
name__iexact=category_param, active=True
|
||||||
|
).first()
|
||||||
if category:
|
if category:
|
||||||
initial_data["category"] = category.pk
|
initial_data["category"] = category.pk
|
||||||
|
|
||||||
@@ -457,7 +555,7 @@ def transaction_pay(request, transaction_id):
|
|||||||
context={"transaction": transaction, **request.GET},
|
context={"transaction": transaction, **request.GET},
|
||||||
)
|
)
|
||||||
response.headers["HX-Trigger"] = (
|
response.headers["HX-Trigger"] = (
|
||||||
f'{"paid" if new_is_paid else "unpaid"}, selective_update'
|
f"{'paid' if new_is_paid else 'unpaid'}, selective_update"
|
||||||
)
|
)
|
||||||
return response
|
return response
|
||||||
|
|
||||||
@@ -537,7 +635,92 @@ def transaction_all_index(request):
|
|||||||
return render(
|
return render(
|
||||||
request,
|
request,
|
||||||
"transactions/pages/transactions.html",
|
"transactions/pages/transactions.html",
|
||||||
{"filter": f, "order": order, "summary_tab": summary_tab},
|
{
|
||||||
|
"filter": f,
|
||||||
|
"filter_is_active": f.has_active_filters,
|
||||||
|
"filter_presets": FilterPreset.objects.filter(owner=request.user),
|
||||||
|
"order": order,
|
||||||
|
"summary_tab": summary_tab,
|
||||||
|
},
|
||||||
|
)
|
||||||
|
|
||||||
|
|
||||||
|
@only_htmx
|
||||||
|
@login_required
|
||||||
|
@require_http_methods(["POST"])
|
||||||
|
def filter_preset_create(request):
|
||||||
|
name = request.POST.get("name", "").strip()
|
||||||
|
if not name or len(name) > 100:
|
||||||
|
return HttpResponse(status=400)
|
||||||
|
|
||||||
|
parameters = {
|
||||||
|
key: request.POST.getlist(key)
|
||||||
|
for key in TransactionsFilter.base_filters
|
||||||
|
if key in request.POST
|
||||||
|
}
|
||||||
|
FilterPreset.objects.create(
|
||||||
|
owner=request.user,
|
||||||
|
name=name,
|
||||||
|
parameters=parameters,
|
||||||
|
)
|
||||||
|
return render(
|
||||||
|
request,
|
||||||
|
"transactions/fragments/filter_presets.html",
|
||||||
|
{"filter_presets": FilterPreset.objects.filter(owner=request.user)},
|
||||||
|
)
|
||||||
|
|
||||||
|
|
||||||
|
@only_htmx
|
||||||
|
@login_required
|
||||||
|
@require_http_methods(["GET"])
|
||||||
|
def filter_preset_apply(request, preset_id):
|
||||||
|
preset = get_object_or_404(FilterPreset, pk=preset_id, owner=request.user)
|
||||||
|
data = QueryDict(mutable=True)
|
||||||
|
for key, values in preset.parameters.items():
|
||||||
|
if key in TransactionsFilter.base_filters:
|
||||||
|
data.setlist(key, values)
|
||||||
|
|
||||||
|
transaction_filter = TransactionsFilter(data)
|
||||||
|
response = render(
|
||||||
|
request,
|
||||||
|
"transactions/fragments/filter_form.html",
|
||||||
|
{
|
||||||
|
"filter": transaction_filter,
|
||||||
|
"filter_is_active": transaction_filter.has_active_filters,
|
||||||
|
"swap_filter_indicator": True,
|
||||||
|
},
|
||||||
|
)
|
||||||
|
response.headers["HX-Trigger-After-Settle"] = "updated"
|
||||||
|
return response
|
||||||
|
|
||||||
|
|
||||||
|
@only_htmx
|
||||||
|
@login_required
|
||||||
|
@require_http_methods(["GET"])
|
||||||
|
def transaction_filter_clear(request):
|
||||||
|
transaction_filter = TransactionsFilter(QueryDict())
|
||||||
|
response = render(
|
||||||
|
request,
|
||||||
|
"transactions/fragments/filter_form.html",
|
||||||
|
{
|
||||||
|
"filter": transaction_filter,
|
||||||
|
"filter_is_active": False,
|
||||||
|
"swap_filter_indicator": True,
|
||||||
|
},
|
||||||
|
)
|
||||||
|
response.headers["HX-Trigger-After-Settle"] = "updated"
|
||||||
|
return response
|
||||||
|
|
||||||
|
|
||||||
|
@only_htmx
|
||||||
|
@login_required
|
||||||
|
@require_http_methods(["POST"])
|
||||||
|
def filter_preset_delete(request, preset_id):
|
||||||
|
get_object_or_404(FilterPreset, pk=preset_id, owner=request.user).delete()
|
||||||
|
return render(
|
||||||
|
request,
|
||||||
|
"transactions/fragments/filter_presets.html",
|
||||||
|
{"filter_presets": FilterPreset.objects.filter(owner=request.user)},
|
||||||
)
|
)
|
||||||
|
|
||||||
|
|
||||||
@@ -552,6 +735,8 @@ def transaction_all_list(request):
|
|||||||
if order != request.session.get("all_transactions_order", "default"):
|
if order != request.session.get("all_transactions_order", "default"):
|
||||||
request.session["all_transactions_order"] = order
|
request.session["all_transactions_order"] = order
|
||||||
|
|
||||||
|
today = timezone.localdate(timezone.now())
|
||||||
|
|
||||||
transactions = Transaction.objects.prefetch_related(
|
transactions = Transaction.objects.prefetch_related(
|
||||||
"account",
|
"account",
|
||||||
"account__group",
|
"account__group",
|
||||||
@@ -565,12 +750,27 @@ def transaction_all_list(request):
|
|||||||
"dca_income_entries",
|
"dca_income_entries",
|
||||||
).all()
|
).all()
|
||||||
|
|
||||||
transactions = default_order(transactions, order=order)
|
|
||||||
|
|
||||||
f = TransactionsFilter(request.GET, queryset=transactions)
|
f = TransactionsFilter(request.GET, queryset=transactions)
|
||||||
|
|
||||||
|
# Late transactions: date < today and is_paid = False (only shown for default ordering on first page)
|
||||||
|
late_transactions = None
|
||||||
page_number = request.GET.get("page", 1)
|
page_number = request.GET.get("page", 1)
|
||||||
paginator = Paginator(f.qs, 100)
|
if order == "default" and str(page_number) == "1":
|
||||||
|
late_transactions = f.qs.filter(
|
||||||
|
date__lt=today,
|
||||||
|
is_paid=False,
|
||||||
|
).order_by("date", "id")
|
||||||
|
# Exclude late transactions from the main paginated list
|
||||||
|
main_transactions = f.qs.exclude(
|
||||||
|
date__lt=today,
|
||||||
|
is_paid=False,
|
||||||
|
)
|
||||||
|
else:
|
||||||
|
main_transactions = f.qs
|
||||||
|
|
||||||
|
main_transactions = default_order(main_transactions, order=order)
|
||||||
|
|
||||||
|
paginator = Paginator(main_transactions, 100)
|
||||||
page_obj = paginator.get_page(page_number)
|
page_obj = paginator.get_page(page_number)
|
||||||
|
|
||||||
return render(
|
return render(
|
||||||
@@ -579,6 +779,7 @@ def transaction_all_list(request):
|
|||||||
{
|
{
|
||||||
"page_obj": page_obj,
|
"page_obj": page_obj,
|
||||||
"paginator": paginator,
|
"paginator": paginator,
|
||||||
|
"late_transactions": late_transactions,
|
||||||
},
|
},
|
||||||
)
|
)
|
||||||
|
|
||||||
|
|||||||
@@ -0,0 +1,75 @@
|
|||||||
|
import logging
|
||||||
|
from allauth.socialaccount.adapter import DefaultSocialAccountAdapter
|
||||||
|
from django.contrib.auth import get_user_model
|
||||||
|
|
||||||
|
User = get_user_model()
|
||||||
|
logger = logging.getLogger(__name__)
|
||||||
|
|
||||||
|
|
||||||
|
class AutoConnectSocialAccountAdapter(DefaultSocialAccountAdapter):
|
||||||
|
"""
|
||||||
|
Custom adapter to automatically connect social accounts to existing users
|
||||||
|
with the same email address.
|
||||||
|
|
||||||
|
SECURITY WARNING:
|
||||||
|
This adapter automatically connects OIDC accounts to existing local accounts
|
||||||
|
based on email matching.
|
||||||
|
|
||||||
|
If your OIDC provider allows unverified emails, this could lead to
|
||||||
|
ACCOUNT TAKEOVER attacks where an attacker creates an OIDC account
|
||||||
|
with someone else's email and gains access to their account.
|
||||||
|
"""
|
||||||
|
|
||||||
|
def pre_social_login(self, request, sociallogin):
|
||||||
|
"""
|
||||||
|
Invoked just after a user successfully authenticates via a
|
||||||
|
social provider, but before the login is actually processed.
|
||||||
|
|
||||||
|
If a user with the same email already exists, connect the social
|
||||||
|
account to that existing user instead of creating a new account.
|
||||||
|
"""
|
||||||
|
# If the social account is already connected to a user, do nothing
|
||||||
|
if sociallogin.is_existing:
|
||||||
|
return
|
||||||
|
|
||||||
|
# Check if we have an email from the social provider
|
||||||
|
if not sociallogin.email_addresses:
|
||||||
|
logger.warning(
|
||||||
|
"OIDC login attempted without email address. "
|
||||||
|
f"Provider: {sociallogin.account.provider}"
|
||||||
|
)
|
||||||
|
return
|
||||||
|
|
||||||
|
# Get the email from the social login
|
||||||
|
email = sociallogin.email_addresses[0].email.lower()
|
||||||
|
|
||||||
|
# Try to find an existing user with this email
|
||||||
|
try:
|
||||||
|
user = User.objects.get(email__iexact=email)
|
||||||
|
|
||||||
|
# Log this connection for security audit trail
|
||||||
|
logger.info(
|
||||||
|
f"Auto-connecting OIDC account to existing user. "
|
||||||
|
f"Email: {email}, Provider: {sociallogin.account.provider}, "
|
||||||
|
f"User ID: {user.id}"
|
||||||
|
)
|
||||||
|
|
||||||
|
# Connect the social account to the existing user
|
||||||
|
sociallogin.connect(request, user)
|
||||||
|
|
||||||
|
except User.DoesNotExist:
|
||||||
|
# No user with this email exists, proceed with normal signup flow
|
||||||
|
logger.debug(
|
||||||
|
f"No existing user found for email {email}. "
|
||||||
|
"Proceeding with new account creation."
|
||||||
|
)
|
||||||
|
pass
|
||||||
|
except User.MultipleObjectsReturned:
|
||||||
|
# Multiple users with the same email (shouldn't happen with unique constraint)
|
||||||
|
logger.error(
|
||||||
|
f"Multiple users found with email {email}. "
|
||||||
|
"This should not happen with unique constraint. "
|
||||||
|
"Blocking auto-connect."
|
||||||
|
)
|
||||||
|
# Let the default behavior handle this
|
||||||
|
pass
|
||||||
+37
-1
@@ -4,13 +4,19 @@ from django.contrib.auth.forms import (
|
|||||||
UserCreationForm,
|
UserCreationForm,
|
||||||
AdminPasswordChangeForm,
|
AdminPasswordChangeForm,
|
||||||
)
|
)
|
||||||
|
from django.utils import timezone
|
||||||
from django.utils.translation import gettext_lazy as _
|
from django.utils.translation import gettext_lazy as _
|
||||||
from django.contrib import admin
|
from django.contrib import admin
|
||||||
from django.contrib.auth.admin import GroupAdmin as BaseGroupAdmin
|
from django.contrib.auth.admin import GroupAdmin as BaseGroupAdmin
|
||||||
from django.contrib.auth.admin import UserAdmin as BaseUserAdmin
|
from django.contrib.auth.admin import UserAdmin as BaseUserAdmin
|
||||||
from django.contrib.auth.models import Group
|
from django.contrib.auth.models import Group
|
||||||
|
|
||||||
from apps.users.models import User, UserSettings
|
from apps.users.models import APIToken, User, UserSettings
|
||||||
|
|
||||||
|
|
||||||
|
@admin.action(description=_("Revoke selected API tokens"))
|
||||||
|
def revoke_api_tokens(modeladmin, request, queryset):
|
||||||
|
queryset.update(revoked_at=timezone.now())
|
||||||
|
|
||||||
admin.site.unregister(Group)
|
admin.site.unregister(Group)
|
||||||
|
|
||||||
@@ -77,3 +83,33 @@ class GroupAdmin(BaseGroupAdmin, ModelAdmin):
|
|||||||
|
|
||||||
|
|
||||||
admin.site.register(UserSettings)
|
admin.site.register(UserSettings)
|
||||||
|
|
||||||
|
|
||||||
|
@admin.register(APIToken)
|
||||||
|
class APITokenAdmin(admin.ModelAdmin):
|
||||||
|
actions = [revoke_api_tokens]
|
||||||
|
list_display = (
|
||||||
|
"name",
|
||||||
|
"user",
|
||||||
|
"token_key",
|
||||||
|
"created_at",
|
||||||
|
"last_used_at",
|
||||||
|
"expires_at",
|
||||||
|
"revoked_at",
|
||||||
|
)
|
||||||
|
search_fields = ("name", "user__email", "token_key")
|
||||||
|
# Never expose the secret hash in the form; it must not be editable.
|
||||||
|
exclude = ("token_hash",)
|
||||||
|
readonly_fields = (
|
||||||
|
"user",
|
||||||
|
"name",
|
||||||
|
"token_key",
|
||||||
|
"created_at",
|
||||||
|
"updated_at",
|
||||||
|
"last_used_at",
|
||||||
|
"expires_at",
|
||||||
|
"revoked_at",
|
||||||
|
)
|
||||||
|
|
||||||
|
def has_add_permission(self, request):
|
||||||
|
return False
|
||||||
|
|||||||
+73
-4
@@ -1,6 +1,11 @@
|
|||||||
|
from datetime import timedelta
|
||||||
|
|
||||||
from apps.common.middleware.thread_local import get_current_user
|
from apps.common.middleware.thread_local import get_current_user
|
||||||
|
from apps.users.models import APIToken
|
||||||
from apps.common.widgets.crispy.submit import NoClassSubmit
|
from apps.common.widgets.crispy.submit import NoClassSubmit
|
||||||
|
from apps.common.widgets.tom_select import TomSelect
|
||||||
from apps.users.models import UserSettings
|
from apps.users.models import UserSettings
|
||||||
|
from apps.accounts.models import Account
|
||||||
from crispy_forms.bootstrap import (
|
from crispy_forms.bootstrap import (
|
||||||
FormActions,
|
FormActions,
|
||||||
)
|
)
|
||||||
@@ -14,6 +19,7 @@ from django.contrib.auth.forms import (
|
|||||||
UsernameField,
|
UsernameField,
|
||||||
)
|
)
|
||||||
from django.db import transaction
|
from django.db import transaction
|
||||||
|
from django.utils import timezone
|
||||||
from django.utils.translation import gettext_lazy as _
|
from django.utils.translation import gettext_lazy as _
|
||||||
|
|
||||||
|
|
||||||
@@ -116,6 +122,15 @@ class UserSettingsForm(forms.ModelForm):
|
|||||||
label=_("Number Format"),
|
label=_("Number Format"),
|
||||||
)
|
)
|
||||||
|
|
||||||
|
default_account = forms.ModelChoiceField(
|
||||||
|
queryset=Account.objects.filter(
|
||||||
|
is_archived=False,
|
||||||
|
),
|
||||||
|
label=_("Default Account"),
|
||||||
|
widget=TomSelect(clear_button=False, group_by="group"),
|
||||||
|
required=False,
|
||||||
|
)
|
||||||
|
|
||||||
class Meta:
|
class Meta:
|
||||||
model = UserSettings
|
model = UserSettings
|
||||||
fields = [
|
fields = [
|
||||||
@@ -125,25 +140,33 @@ class UserSettingsForm(forms.ModelForm):
|
|||||||
"date_format",
|
"date_format",
|
||||||
"datetime_format",
|
"datetime_format",
|
||||||
"number_format",
|
"number_format",
|
||||||
"volume",
|
"default_account",
|
||||||
]
|
]
|
||||||
|
widgets = {
|
||||||
|
"default_account": TomSelect(clear_button=False, group_by="group"),
|
||||||
|
}
|
||||||
|
|
||||||
def __init__(self, *args, **kwargs):
|
def __init__(self, *args, **kwargs):
|
||||||
super().__init__(*args, **kwargs)
|
super().__init__(*args, **kwargs)
|
||||||
|
|
||||||
|
self.fields["default_account"].queryset = Account.objects.filter(
|
||||||
|
is_archived=False,
|
||||||
|
)
|
||||||
|
|
||||||
self.helper = FormHelper()
|
self.helper = FormHelper()
|
||||||
self.helper.form_tag = False
|
self.helper.form_tag = False
|
||||||
self.helper.form_method = "post"
|
self.helper.form_method = "post"
|
||||||
self.helper.layout = Layout(
|
self.helper.layout = Layout(
|
||||||
"language",
|
"language",
|
||||||
"timezone",
|
"timezone",
|
||||||
HTML("<hr />"),
|
HTML('<hr class="hr my-3" />'),
|
||||||
"date_format",
|
"date_format",
|
||||||
"datetime_format",
|
"datetime_format",
|
||||||
"number_format",
|
"number_format",
|
||||||
HTML("<hr />"),
|
HTML('<hr class="hr my-3" />'),
|
||||||
"start_page",
|
"start_page",
|
||||||
HTML("<hr />"),
|
"default_account",
|
||||||
|
HTML('<hr class="hr my-3" />'),
|
||||||
"volume",
|
"volume",
|
||||||
FormActions(
|
FormActions(
|
||||||
NoClassSubmit("submit", _("Save"), css_class="btn btn-primary"),
|
NoClassSubmit("submit", _("Save"), css_class="btn btn-primary"),
|
||||||
@@ -407,3 +430,49 @@ class UserAddForm(UserCreationForm):
|
|||||||
if commit:
|
if commit:
|
||||||
user.save()
|
user.save()
|
||||||
return user
|
return user
|
||||||
|
|
||||||
|
|
||||||
|
class APITokenCreateForm(forms.Form):
|
||||||
|
name = forms.CharField(
|
||||||
|
max_length=255,
|
||||||
|
label=_("Token name"),
|
||||||
|
help_text=_(
|
||||||
|
"Use a descriptive name such as n8n, Home Assistant, or backup job."
|
||||||
|
),
|
||||||
|
)
|
||||||
|
expires_in_days = forms.IntegerField(
|
||||||
|
required=False,
|
||||||
|
min_value=1,
|
||||||
|
label=_("Expires in days"),
|
||||||
|
help_text=_("Leave empty for a non-expiring token."),
|
||||||
|
)
|
||||||
|
|
||||||
|
def __init__(self, *args, **kwargs):
|
||||||
|
super().__init__(*args, **kwargs)
|
||||||
|
|
||||||
|
self.helper = FormHelper()
|
||||||
|
self.helper.form_tag = False
|
||||||
|
self.helper.form_method = "post"
|
||||||
|
self.helper.layout = Layout(
|
||||||
|
"name",
|
||||||
|
"expires_in_days",
|
||||||
|
FormActions(
|
||||||
|
NoClassSubmit(
|
||||||
|
"submit",
|
||||||
|
_("Create token"),
|
||||||
|
css_class="btn btn-primary",
|
||||||
|
),
|
||||||
|
),
|
||||||
|
)
|
||||||
|
|
||||||
|
def save(self, user):
|
||||||
|
expires_in_days = self.cleaned_data.get("expires_in_days")
|
||||||
|
expires_at = None
|
||||||
|
if expires_in_days:
|
||||||
|
expires_at = timezone.now() + timedelta(days=expires_in_days)
|
||||||
|
|
||||||
|
return APIToken.objects.create_token(
|
||||||
|
user=user,
|
||||||
|
name=self.cleaned_data["name"],
|
||||||
|
expires_at=expires_at,
|
||||||
|
)
|
||||||
|
|||||||
@@ -0,0 +1,25 @@
|
|||||||
|
# Generated by Django 5.2.9 on 2026-02-15 21:35
|
||||||
|
|
||||||
|
import django.db.models.deletion
|
||||||
|
from django.db import migrations, models
|
||||||
|
|
||||||
|
|
||||||
|
class Migration(migrations.Migration):
|
||||||
|
dependencies = [
|
||||||
|
("accounts", "0016_account_untracked_by"),
|
||||||
|
("users", "0023_alter_usersettings_timezone"),
|
||||||
|
]
|
||||||
|
|
||||||
|
operations = [
|
||||||
|
migrations.AddField(
|
||||||
|
model_name="usersettings",
|
||||||
|
name="default_account",
|
||||||
|
field=models.ForeignKey(
|
||||||
|
blank=True,
|
||||||
|
null=True,
|
||||||
|
on_delete=django.db.models.deletion.SET_NULL,
|
||||||
|
to="accounts.account",
|
||||||
|
verbose_name="Default account",
|
||||||
|
),
|
||||||
|
),
|
||||||
|
]
|
||||||
@@ -0,0 +1,20 @@
|
|||||||
|
# Generated by Django 5.2.9 on 2026-02-16 01:32
|
||||||
|
|
||||||
|
import django.db.models.deletion
|
||||||
|
from django.db import migrations, models
|
||||||
|
|
||||||
|
|
||||||
|
class Migration(migrations.Migration):
|
||||||
|
|
||||||
|
dependencies = [
|
||||||
|
('accounts', '0016_account_untracked_by'),
|
||||||
|
('users', '0024_usersettings_default_account'),
|
||||||
|
]
|
||||||
|
|
||||||
|
operations = [
|
||||||
|
migrations.AlterField(
|
||||||
|
model_name='usersettings',
|
||||||
|
name='default_account',
|
||||||
|
field=models.ForeignKey(blank=True, help_text='Selects the account by default when creating new transactions', null=True, on_delete=django.db.models.deletion.SET_NULL, to='accounts.account', verbose_name='Default account'),
|
||||||
|
),
|
||||||
|
]
|
||||||
@@ -0,0 +1,36 @@
|
|||||||
|
# Generated by Django 5.2.15 on 2026-06-24 09:21
|
||||||
|
|
||||||
|
import django.db.models.deletion
|
||||||
|
from django.conf import settings
|
||||||
|
from django.db import migrations, models
|
||||||
|
|
||||||
|
|
||||||
|
class Migration(migrations.Migration):
|
||||||
|
|
||||||
|
dependencies = [
|
||||||
|
('users', '0025_alter_usersettings_default_account'),
|
||||||
|
]
|
||||||
|
|
||||||
|
operations = [
|
||||||
|
migrations.CreateModel(
|
||||||
|
name='APIToken',
|
||||||
|
fields=[
|
||||||
|
('id', models.BigAutoField(auto_created=True, primary_key=True, serialize=False, verbose_name='ID')),
|
||||||
|
('name', models.CharField(max_length=255, verbose_name='Name')),
|
||||||
|
('token_key', models.CharField(db_index=True, max_length=16, unique=True, verbose_name='Token key')),
|
||||||
|
('token_hash', models.CharField(max_length=255, verbose_name='Token hash')),
|
||||||
|
('last_used_at', models.DateTimeField(blank=True, null=True, verbose_name='Last used at')),
|
||||||
|
('expires_at', models.DateTimeField(blank=True, null=True, verbose_name='Expires at')),
|
||||||
|
('revoked_at', models.DateTimeField(blank=True, null=True, verbose_name='Revoked at')),
|
||||||
|
('created_at', models.DateTimeField(auto_now_add=True, verbose_name='Created at')),
|
||||||
|
('updated_at', models.DateTimeField(auto_now=True, verbose_name='Updated at')),
|
||||||
|
('user', models.ForeignKey(on_delete=django.db.models.deletion.CASCADE, related_name='api_tokens', to=settings.AUTH_USER_MODEL, verbose_name='User')),
|
||||||
|
],
|
||||||
|
options={
|
||||||
|
'verbose_name': 'API token',
|
||||||
|
'verbose_name_plural': 'API tokens',
|
||||||
|
'ordering': ['-created_at'],
|
||||||
|
'indexes': [models.Index(fields=['user', 'revoked_at'], name='users_apito_user_id_73edec_idx'), models.Index(fields=['expires_at'], name='users_apito_expires_2b737c_idx')],
|
||||||
|
},
|
||||||
|
),
|
||||||
|
]
|
||||||
File diff suppressed because one or more lines are too long
+127
-2
@@ -1,9 +1,14 @@
|
|||||||
|
import hashlib
|
||||||
|
import hmac
|
||||||
|
import secrets
|
||||||
|
|
||||||
import pytz
|
import pytz
|
||||||
from django.conf import settings
|
from django.conf import settings
|
||||||
from django.contrib.auth import get_user_model
|
from django.contrib.auth import get_user_model
|
||||||
from django.contrib.auth.models import AbstractUser, Group
|
from django.contrib.auth.models import AbstractUser, Group
|
||||||
from django.core.validators import MaxValueValidator, MinValueValidator
|
from django.core.validators import MaxValueValidator, MinValueValidator
|
||||||
from django.db import models
|
from django.db import IntegrityError, models, transaction
|
||||||
|
from django.utils import timezone
|
||||||
from django.utils.translation import gettext_lazy as _
|
from django.utils.translation import gettext_lazy as _
|
||||||
|
|
||||||
from apps.users.managers import UserManager
|
from apps.users.managers import UserManager
|
||||||
@@ -410,7 +415,7 @@ timezones = [
|
|||||||
("Pacific/Galapagos", "Pacific/Galapagos"),
|
("Pacific/Galapagos", "Pacific/Galapagos"),
|
||||||
("Pacific/Gambier", "Pacific/Gambier"),
|
("Pacific/Gambier", "Pacific/Gambier"),
|
||||||
("Pacific/Guadalcanal", "Pacific/Guadalcanal"),
|
("Pacific/Guadalcanal", "Pacific/Guadalcanal"),
|
||||||
("P2025-06-29T01:43:14.671389745Z acific/Guam", "Pacific/Guam"),
|
("Pacific/Guam", "Pacific/Guam"),
|
||||||
("Pacific/Honolulu", "Pacific/Honolulu"),
|
("Pacific/Honolulu", "Pacific/Honolulu"),
|
||||||
("Pacific/Kanton", "Pacific/Kanton"),
|
("Pacific/Kanton", "Pacific/Kanton"),
|
||||||
("Pacific/Kiritimati", "Pacific/Kiritimati"),
|
("Pacific/Kiritimati", "Pacific/Kiritimati"),
|
||||||
@@ -510,9 +515,129 @@ class UserSettings(models.Model):
|
|||||||
default=StartPage.MONTHLY,
|
default=StartPage.MONTHLY,
|
||||||
verbose_name=_("Start page"),
|
verbose_name=_("Start page"),
|
||||||
)
|
)
|
||||||
|
default_account = models.ForeignKey(
|
||||||
|
"accounts.Account",
|
||||||
|
on_delete=models.SET_NULL,
|
||||||
|
verbose_name=_("Default account"),
|
||||||
|
help_text=_("Selects the account by default when creating new transactions"),
|
||||||
|
blank=True,
|
||||||
|
null=True,
|
||||||
|
)
|
||||||
|
|
||||||
def __str__(self):
|
def __str__(self):
|
||||||
return f"{self.user.email}'s settings"
|
return f"{self.user.email}'s settings"
|
||||||
|
|
||||||
def clean(self):
|
def clean(self):
|
||||||
super().clean()
|
super().clean()
|
||||||
|
|
||||||
|
|
||||||
|
class APITokenManager(models.Manager):
|
||||||
|
def create_token(self, *, user, name: str, expires_at=None):
|
||||||
|
token_secret = secrets.token_urlsafe(32)
|
||||||
|
token_hash = self.model.hash_secret(token_secret)
|
||||||
|
|
||||||
|
# token_key is unique; the pre-check in generate_token_key still leaves a
|
||||||
|
# tiny race window under concurrency, so retry on the unique-constraint
|
||||||
|
# violation with a fresh key instead of failing the request.
|
||||||
|
last_error = None
|
||||||
|
for _ in range(5):
|
||||||
|
token = self.model(
|
||||||
|
user=user,
|
||||||
|
name=name,
|
||||||
|
token_key=self.model.generate_token_key(),
|
||||||
|
token_hash=token_hash,
|
||||||
|
expires_at=expires_at,
|
||||||
|
)
|
||||||
|
token.full_clean()
|
||||||
|
try:
|
||||||
|
with transaction.atomic():
|
||||||
|
token.save()
|
||||||
|
except IntegrityError as exc:
|
||||||
|
last_error = exc
|
||||||
|
continue
|
||||||
|
return token, token.build_raw_token(token_secret)
|
||||||
|
|
||||||
|
raise last_error
|
||||||
|
|
||||||
|
|
||||||
|
class APIToken(models.Model):
|
||||||
|
TOKEN_PREFIX = "wygiwyh_pat_"
|
||||||
|
|
||||||
|
user = models.ForeignKey(
|
||||||
|
settings.AUTH_USER_MODEL,
|
||||||
|
on_delete=models.CASCADE,
|
||||||
|
related_name="api_tokens",
|
||||||
|
verbose_name=_("User"),
|
||||||
|
)
|
||||||
|
name = models.CharField(max_length=255, verbose_name=_("Name"))
|
||||||
|
token_key = models.CharField(
|
||||||
|
max_length=16,
|
||||||
|
unique=True,
|
||||||
|
db_index=True,
|
||||||
|
verbose_name=_("Token key"),
|
||||||
|
)
|
||||||
|
token_hash = models.CharField(max_length=255, verbose_name=_("Token hash"))
|
||||||
|
last_used_at = models.DateTimeField(
|
||||||
|
null=True,
|
||||||
|
blank=True,
|
||||||
|
verbose_name=_("Last used at"),
|
||||||
|
)
|
||||||
|
expires_at = models.DateTimeField(
|
||||||
|
null=True,
|
||||||
|
blank=True,
|
||||||
|
verbose_name=_("Expires at"),
|
||||||
|
)
|
||||||
|
revoked_at = models.DateTimeField(
|
||||||
|
null=True,
|
||||||
|
blank=True,
|
||||||
|
verbose_name=_("Revoked at"),
|
||||||
|
)
|
||||||
|
created_at = models.DateTimeField(auto_now_add=True, verbose_name=_("Created at"))
|
||||||
|
updated_at = models.DateTimeField(auto_now=True, verbose_name=_("Updated at"))
|
||||||
|
|
||||||
|
objects = APITokenManager()
|
||||||
|
|
||||||
|
class Meta:
|
||||||
|
indexes = [
|
||||||
|
models.Index(fields=["user", "revoked_at"]),
|
||||||
|
models.Index(fields=["expires_at"]),
|
||||||
|
]
|
||||||
|
ordering = ["-created_at"]
|
||||||
|
verbose_name = _("API token")
|
||||||
|
verbose_name_plural = _("API tokens")
|
||||||
|
|
||||||
|
def __str__(self):
|
||||||
|
return f"{self.user} / {self.name}"
|
||||||
|
|
||||||
|
@classmethod
|
||||||
|
def generate_token_key(cls) -> str:
|
||||||
|
while True:
|
||||||
|
candidate = secrets.token_hex(8)
|
||||||
|
if not cls.objects.filter(token_key=candidate).exists():
|
||||||
|
return candidate
|
||||||
|
|
||||||
|
@classmethod
|
||||||
|
def parse_raw_token(cls, raw_token: str):
|
||||||
|
if not raw_token.startswith(cls.TOKEN_PREFIX):
|
||||||
|
raise ValueError("Token is missing the expected prefix.")
|
||||||
|
|
||||||
|
payload = raw_token.removeprefix(cls.TOKEN_PREFIX)
|
||||||
|
token_key, separator, token_secret = payload.partition(".")
|
||||||
|
if not separator or not token_key or not token_secret:
|
||||||
|
raise ValueError("Token is malformed.")
|
||||||
|
return token_key, token_secret
|
||||||
|
|
||||||
|
def build_raw_token(self, token_secret: str) -> str:
|
||||||
|
return f"{self.TOKEN_PREFIX}{self.token_key}.{token_secret}"
|
||||||
|
|
||||||
|
@staticmethod
|
||||||
|
def hash_secret(token_secret: str) -> str:
|
||||||
|
# The secret is a 256-bit random value (secrets.token_urlsafe(32)), so a
|
||||||
|
# single SHA-256 is sufficient and avoids a slow KDF on every request.
|
||||||
|
return hashlib.sha256(token_secret.encode("utf-8")).hexdigest()
|
||||||
|
|
||||||
|
def check_secret(self, raw_secret: str) -> bool:
|
||||||
|
return hmac.compare_digest(self.token_hash, self.hash_secret(raw_secret))
|
||||||
|
|
||||||
|
def is_expired(self) -> bool:
|
||||||
|
return self.expires_at is not None and self.expires_at <= timezone.now()
|
||||||
|
|||||||
@@ -0,0 +1,73 @@
|
|||||||
|
from django.contrib.auth import get_user_model
|
||||||
|
from django.test import TestCase
|
||||||
|
from django.urls import reverse
|
||||||
|
from django.utils import timezone
|
||||||
|
|
||||||
|
from apps.users.models import APIToken
|
||||||
|
|
||||||
|
|
||||||
|
class UserAPITokenViewsTests(TestCase):
|
||||||
|
def setUp(self):
|
||||||
|
self.user = get_user_model().objects.create_user(
|
||||||
|
email="user@example.com",
|
||||||
|
password="test-password",
|
||||||
|
)
|
||||||
|
self.client.force_login(self.user)
|
||||||
|
self.htmx_headers = {"HTTP_HX_REQUEST": "true"}
|
||||||
|
|
||||||
|
def test_user_settings_renders_api_token_section(self):
|
||||||
|
response = self.client.get(reverse("user_settings"), **self.htmx_headers)
|
||||||
|
|
||||||
|
self.assertContains(response, "API Tokens")
|
||||||
|
self.assertContains(response, reverse("user_api_token_add"))
|
||||||
|
|
||||||
|
def test_can_create_api_token_from_ui(self):
|
||||||
|
response = self.client.post(
|
||||||
|
reverse("user_api_token_add"),
|
||||||
|
{"name": "n8n", "expires_in_days": "30"},
|
||||||
|
**self.htmx_headers,
|
||||||
|
)
|
||||||
|
|
||||||
|
self.assertEqual(response.status_code, 200)
|
||||||
|
self.assertContains(response, "Copy this token now")
|
||||||
|
self.assertEqual(APIToken.objects.filter(user=self.user, name="n8n").count(), 1)
|
||||||
|
|
||||||
|
def test_can_revoke_own_api_token(self):
|
||||||
|
token, _ = APIToken.objects.create_token(user=self.user, name="n8n")
|
||||||
|
|
||||||
|
response = self.client.delete(
|
||||||
|
reverse("user_api_token_revoke", kwargs={"token_id": token.id}),
|
||||||
|
**self.htmx_headers,
|
||||||
|
)
|
||||||
|
|
||||||
|
self.assertEqual(response.status_code, 200)
|
||||||
|
token.refresh_from_db()
|
||||||
|
self.assertIsNotNone(token.revoked_at)
|
||||||
|
self.assertContains(response, "Revoked")
|
||||||
|
|
||||||
|
def test_can_delete_revoked_api_token(self):
|
||||||
|
token, _ = APIToken.objects.create_token(user=self.user, name="n8n")
|
||||||
|
token.revoked_at = timezone.now()
|
||||||
|
token.save(update_fields=["revoked_at"])
|
||||||
|
|
||||||
|
response = self.client.delete(
|
||||||
|
reverse("user_api_token_delete", kwargs={"token_id": token.id}),
|
||||||
|
**self.htmx_headers,
|
||||||
|
)
|
||||||
|
|
||||||
|
self.assertEqual(response.status_code, 200)
|
||||||
|
self.assertFalse(APIToken.objects.filter(id=token.id).exists())
|
||||||
|
|
||||||
|
def test_cannot_delete_other_users_api_token(self):
|
||||||
|
other = get_user_model().objects.create_user(
|
||||||
|
email="other@example.com", password="test-password"
|
||||||
|
)
|
||||||
|
token, _ = APIToken.objects.create_token(user=other, name="theirs")
|
||||||
|
|
||||||
|
response = self.client.delete(
|
||||||
|
reverse("user_api_token_delete", kwargs={"token_id": token.id}),
|
||||||
|
**self.htmx_headers,
|
||||||
|
)
|
||||||
|
|
||||||
|
self.assertEqual(response.status_code, 404)
|
||||||
|
self.assertTrue(APIToken.objects.filter(id=token.id).exists())
|
||||||
@@ -32,6 +32,21 @@ urlpatterns = [
|
|||||||
views.update_settings,
|
views.update_settings,
|
||||||
name="user_settings",
|
name="user_settings",
|
||||||
),
|
),
|
||||||
|
path(
|
||||||
|
"user/api-tokens/add/",
|
||||||
|
views.api_token_add,
|
||||||
|
name="user_api_token_add",
|
||||||
|
),
|
||||||
|
path(
|
||||||
|
"user/api-tokens/<int:token_id>/revoke/",
|
||||||
|
views.api_token_revoke,
|
||||||
|
name="user_api_token_revoke",
|
||||||
|
),
|
||||||
|
path(
|
||||||
|
"user/api-tokens/<int:token_id>/delete/",
|
||||||
|
views.api_token_delete,
|
||||||
|
name="user_api_token_delete",
|
||||||
|
),
|
||||||
path(
|
path(
|
||||||
"users/",
|
"users/",
|
||||||
views.users_index,
|
views.users_index,
|
||||||
|
|||||||
+66
-2
@@ -2,12 +2,13 @@ from apps.common.decorators.demo import disabled_on_demo
|
|||||||
from apps.common.decorators.htmx import only_htmx
|
from apps.common.decorators.htmx import only_htmx
|
||||||
from apps.common.decorators.user import htmx_login_required, is_superuser
|
from apps.common.decorators.user import htmx_login_required, is_superuser
|
||||||
from apps.users.forms import (
|
from apps.users.forms import (
|
||||||
|
APITokenCreateForm,
|
||||||
LoginForm,
|
LoginForm,
|
||||||
UserAddForm,
|
UserAddForm,
|
||||||
UserSettingsForm,
|
UserSettingsForm,
|
||||||
UserUpdateForm,
|
UserUpdateForm,
|
||||||
)
|
)
|
||||||
from apps.users.models import UserSettings
|
from apps.users.models import APIToken, UserSettings
|
||||||
from django.contrib import messages
|
from django.contrib import messages
|
||||||
from django.contrib.auth import get_user_model, logout
|
from django.contrib.auth import get_user_model, logout
|
||||||
from django.contrib.auth.decorators import login_required
|
from django.contrib.auth.decorators import login_required
|
||||||
@@ -18,6 +19,7 @@ from django.core.exceptions import PermissionDenied
|
|||||||
from django.http import HttpResponse
|
from django.http import HttpResponse
|
||||||
from django.shortcuts import get_object_or_404, redirect, render
|
from django.shortcuts import get_object_or_404, redirect, render
|
||||||
from django.urls import reverse
|
from django.urls import reverse
|
||||||
|
from django.utils import timezone
|
||||||
from django.utils.translation import gettext_lazy as _
|
from django.utils.translation import gettext_lazy as _
|
||||||
from django.views.decorators.http import require_http_methods
|
from django.views.decorators.http import require_http_methods
|
||||||
|
|
||||||
@@ -112,7 +114,69 @@ def update_settings(request):
|
|||||||
else:
|
else:
|
||||||
form = UserSettingsForm(instance=user_settings)
|
form = UserSettingsForm(instance=user_settings)
|
||||||
|
|
||||||
return render(request, "users/fragments/user_settings.html", {"form": form})
|
return render(
|
||||||
|
request,
|
||||||
|
"users/fragments/user_settings.html",
|
||||||
|
{
|
||||||
|
"form": form,
|
||||||
|
"api_token_form": APITokenCreateForm(),
|
||||||
|
"api_tokens": request.user.api_tokens.all(),
|
||||||
|
},
|
||||||
|
)
|
||||||
|
|
||||||
|
|
||||||
|
def _render_api_tokens(request, *, form=None, raw_token=None):
|
||||||
|
return render(
|
||||||
|
request,
|
||||||
|
"users/fragments/api_tokens.html",
|
||||||
|
{
|
||||||
|
"api_token_form": form or APITokenCreateForm(),
|
||||||
|
"api_tokens": request.user.api_tokens.all(),
|
||||||
|
"raw_token": raw_token,
|
||||||
|
},
|
||||||
|
)
|
||||||
|
|
||||||
|
|
||||||
|
@only_htmx
|
||||||
|
@htmx_login_required
|
||||||
|
@disabled_on_demo
|
||||||
|
@require_http_methods(["POST"])
|
||||||
|
def api_token_add(request):
|
||||||
|
form = APITokenCreateForm(request.POST)
|
||||||
|
if form.is_valid():
|
||||||
|
_token, raw_token = form.save(user=request.user)
|
||||||
|
messages.success(request, _("API token created successfully"))
|
||||||
|
return _render_api_tokens(
|
||||||
|
request,
|
||||||
|
form=APITokenCreateForm(),
|
||||||
|
raw_token=raw_token,
|
||||||
|
)
|
||||||
|
|
||||||
|
return _render_api_tokens(request, form=form)
|
||||||
|
|
||||||
|
|
||||||
|
@only_htmx
|
||||||
|
@htmx_login_required
|
||||||
|
@disabled_on_demo
|
||||||
|
@require_http_methods(["DELETE"])
|
||||||
|
def api_token_revoke(request, token_id):
|
||||||
|
token = get_object_or_404(APIToken, id=token_id, user=request.user)
|
||||||
|
if token.revoked_at is None:
|
||||||
|
token.revoked_at = timezone.now()
|
||||||
|
token.save(update_fields=["revoked_at"])
|
||||||
|
messages.success(request, _("API token revoked successfully"))
|
||||||
|
return _render_api_tokens(request)
|
||||||
|
|
||||||
|
|
||||||
|
@only_htmx
|
||||||
|
@htmx_login_required
|
||||||
|
@disabled_on_demo
|
||||||
|
@require_http_methods(["DELETE"])
|
||||||
|
def api_token_delete(request, token_id):
|
||||||
|
token = get_object_or_404(APIToken, id=token_id, user=request.user)
|
||||||
|
token.delete()
|
||||||
|
messages.success(request, _("API token deleted successfully"))
|
||||||
|
return _render_api_tokens(request)
|
||||||
|
|
||||||
|
|
||||||
@only_htmx
|
@only_htmx
|
||||||
|
|||||||
Some files were not shown because too many files have changed in this diff Show More
Reference in New Issue
Block a user